from __future__ import division
from pylab import *
from mpl_toolkits.mplot3d import Axes3D

def euler(f, t0, y0, h, N):
    t = t0 + arange(N+1)*h
    y = zeros((N+1, size(y0)))
    y[0] = y0
    for n in range(N):
        y[n+1] = y[n] + h*f(t[n], y[n])
    return y

def etrap(f, t0, y0, h, N):
    t = t0 + arange(N+1)*h
    y = zeros((N+1, size(y0)))
    y[0] = y0
    for n in range(N):
        xi1 = y[n]
        f1 = f(t[n], xi1)

        xi2 = y[n] + h*f1
        f2 = f(t[n+1], xi2)

        y[n+1] = y[n] + 0.5*h*(f1 + f2)
    return y

def errorPlot(method):
    def f(t,y):
        return y
    err = zeros(10)
    H = 0.5**arange(10)
    for i, h in enumerate(H):
        N = int(1/h)
        y = method(f, 0, 1, h, N)
        err[i] = abs(y[N] - e)
    loglog(H,err)
    xlabel("h")
    ylabel("error")

def lorenzPlot():
    y = rk4(fLorenz, 0, array([0,2,20]), .01, 10000)
    fig = figure()
    ax = Axes3D(fig)
    ax.plot(*y.T)
