#!/usr/bin/env python
# coding: utf-8

# In[11]:


#Euler method for approximating an ODE solution with n equal steps
#where y'(x) = f(x, y) and y(a) = ya
from numpy import empty
def euler(f, a, b, n, ya):
    h = (b-a) / n
    x = empty(n+1)
    y = empty(n+1)
    y[0] = ya
    for i in range(n):
        x[i] = a + i*h
        y[i+1] = y[i] + h*f(x[i], y[i])
    x[n] = b    
    return x, y
    


# In[9]:


f = lambda x, y: x - y
a = 0
b = 2
ya = 1
n = 3


# In[12]:


x, y = euler(f, a, b, n, ya)
x, y


# In[13]:


x, y = euler(f, a, b, 300, ya)
x[-1], y[-1]


# In[14]:


from matplotlib.pyplot import plot
plot(x, y)


# In[15]:


#Heun's method for approximating an ODE solution with n equal steps
#where y'(x) = f(x, y) and y(a) = ya
from numpy import empty
def heun(f, a, b, n, ya):
    h = (b-a) / n
    x = empty(n+1)
    y = empty(n+1)
    y[0] = ya
    for i in range(n):
        x[i] = a + i*h
        K1 = h*f(x[i], y[i])
        K2 = h*f(x[i]+h, y[i]+K1)
        y[i+1] = y[i] + (K1 + K2) / 2
    x[n] = b    
    return x, y


# In[16]:


x, y = heun(f, a, b, 3, ya)
x, y


# In[17]:


x, y = heun(f, a, b, 300, ya)
x[-1], y[-1]


# In[18]:


#RK4 method for approximating an ODE solution with n equal steps
#where y'(x) = f(x, y) and y(a) = ya
from numpy import empty
def rk4(f, a, b, n, ya):
    h = (b-a) / n
    x = empty(n+1)
    y = empty(n+1)
    y[0] = ya
    for i in range(n):
        x[i] = a + i*h
        K1 = h*f(x[i], y[i])
        K2 = h*f(x[i]+h/2, y[i]+K1/2)
        K3 = h*f(x[i]+h/2, y[i]+K2/2)
        K4 = h*f(x[i]+h, y[i]+K3)          
        y[i+1] = y[i] + (K1 + 2*K2 + 2*K3 + K4) / 6
    x[n] = b    
    return x, y


# In[19]:


x, y = rk4(f, a, b, 3, ya)
x, y


# In[20]:


x, y = rk4(f, a, b, 300, ya)
x[-1], y[-1]


# In[43]:


#RK4 method for approximating an ODE solution with n equal steps
#where y'(x) = f(x, y) and y(a) = ya [where ya is a vector and f is vector-valued]
from numpy import empty, size
def rk4_v(f, a, b, n, ya):
    h = (b-a) / n
    m = ya.size
    x = empty(n+1)
    y = empty((m, n+1))
    y[:, 0] = ya
    for i in range(n):
        x[i] = a + i*h
        K1 = h*f(x[i], y[:, i])
        K2 = h*f(x[i]+h/2, y[:, i]+K1/2)
        K3 = h*f(x[i]+h/2, y[:, i]+K2/2)
        K4 = h*f(x[i]+h, y[:, i]+K3)          
        y[:, i+1] = y[:, i] + (K1 + 2*K2 + 2*K3 + K4) / 6
    x[n] = b    
    return x, y


# In[44]:


#sine wave example
from numpy import array
f = lambda x, z: array([z[1], -z[0]])
ya = array([1, 0])
x, y = rk4_v(f, 0, 20, 200, ya)


# In[45]:


y.shape


# In[46]:


plot(x, y[0, :])


# In[35]:


#built-in IVP numerical solution function
from scipy.integrate import solve_ivp
sol = solve_ivp(f, [0, 20], ya, dense_output=True)


# In[38]:


sol.y.shape


# In[49]:


#estimate the values of z(x) at more points in the interval
from numpy import linspace

t_new = linspace(0, 20, 300)
y_new = sol.sol(t_new)


# In[50]:


plot(sol.t, sol.y[0, :], 's', t_new, y_new[0, :])


# In[ ]:




