import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from sklearn.preprocessing import MinMaxScaler
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import LSTM, Dense

# Load data
data = pd.read_csv('/content/crypto_500_records.csv')
data['Date'] = pd.to_datetime(data['Date'])
data.set_index('Date', inplace=True)

symbols = ['BTC-USD', 'ETH-USD']
symbols = [s for s in symbols if s in data.columns]

print("Starting Date:", data.index[0].strftime('%Y-%m-%d'))
print("Ending Date:", data.index[-1].strftime('%Y-%m-%d'))
print("Using:", symbols)

# Scale data
scalers = {}
scaled = pd.DataFrame(index=data.index)

for s in symbols:
    scalers[s] = MinMaxScaler()
    scaled[s] = scalers[s].fit_transform(
        data[[s]]
    ).ravel()

# Create sequences
seq_len = 30

def sequences(symbol):
    X, y = [], []
    for i in range(len(scaled) - seq_len):
        X.append(scaled.iloc[i:i+seq_len].values)
        y.append(scaled[symbol].iloc[i+seq_len])
    return np.array(X), np.array(y)

# Train and predict
for symbol in symbols:

    print(f"\n--- Training {symbol} Model ---")

    X, y = sequences(symbol)

    split = int(0.8 * len(X))

    X_train, X_test = X[:split], X[split:]
    y_train, y_test = y[:split], y[split:]

    model = Sequential([
        LSTM(50, return_sequences=True,
             input_shape=(seq_len, len(symbols))),
        LSTM(50),
        Dense(1)
    ])

    model.compile(optimizer='adam',
                  loss='mean_squared_error')

    model.fit(X_train, y_train,
              epochs=10,
              batch_size=32,
              verbose=0)

    # Prediction
    prediction = model.predict(X_test, verbose=0)

    prediction = scalers[symbol].inverse_transform(prediction)
    actual = scalers[symbol].inverse_transform(
        y_test.reshape(-1, 1)
    )

    # Plot
    dates = data.index[split + seq_len:]

    plt.figure(figsize=(14, 7))
    plt.plot(dates, actual, label=f'Actual ({symbol})')
    plt.plot(dates, prediction, label=f'Predicted ({symbol})')

    plt.xlabel('Date')
    plt.ylabel('Price')
    plt.title(f'{symbol} Prices - Actual vs Predicted')
    plt.legend()
    plt.xticks(rotation=45)
    plt.grid(True)
    plt.tight_layout()
    plt.show()