import numpy as np
from sklearn.cluster import KMeans
from sklearn.neural_network import MLPClassifier
import pylab as plt
import yfinance as yf

print("Downloading BTC data...")
btc = yf.download("BTC-USD", period="2y", interval="1d")

# Use Close price, fill NaN
price = btc['Close'].dropna().values.flatten()
if price.size == 0:
    raise RuntimeError(
        "No BTC data was downloaded. Check internet access, yfinance availability, "
        "or provide a cached data source before running this forecast."
    )
print(f"BTC data loaded: {len(price)} days")

# Train on log returns, then reconstruct price from the last known close.
log_price = np.log(price)
returns = np.diff(log_price)
PAST_STATES = 5
f = MLPClassifier(random_state=42)

N = len(returns) # Equal NN for convergence 0.0 MSE
NN = len(returns)
classes = np.arange(N)

g = KMeans(n_clusters=N)

def predict(steps):
    out = []
    x = series[-PAST_STATES:].reshape(1, -1)
    for _ in range(steps):
        cluster = f.predict(x)[0]
        next_state = M[cluster].copy()
        out.append(next_state)
        x = np.concatenate([x[0, next_state.size:], next_state])[None, :]
    return np.array(out)

#t = np.linspace(0,10,NN)
#series = np.sin(3*t) - np.sin(5*t) + np.sin(7*t)
#series = (3*t)/(np.sin(3*t) + 1e-8)
series = np.hstack([returns[:, None], 0.1 * np.random.randn(NN, 9)])

g.fit(series)
M = g.cluster_centers_

i=0
while True:
    X = []
    yt = []
    for idx in np.arange(PAST_STATES,NN):
        yt.append(series[idx:idx+1])
        X.append(series[idx-PAST_STATES:idx].reshape(1, -1))
    yt_ = np.array(yt).squeeze(1)
    X_ = np.array(X).squeeze(1)
    p = g.predict(yt_)
    f.partial_fit(X_, p, classes=classes)
    err = yt_ - M[f.predict(X_)]
    print(i,np.sum(err**2))
    if i==130:break
    i+=1

plt.plot(price)
predicted_returns = predict(400)[:, 0]
predict400 = np.concatenate([price, price[-1] * np.exp(np.cumsum(predicted_returns))])
plt.plot(predict400)
plt.show()
