import torch
import torch.nn as nn
import torch.nn.functional as F
from sklearn.datasets import fetch_openml
from sklearn.model_selection import train_test_split
from ionic_ml_simulator import IonicMLSimulator

# Load MNIST subset
X, y = fetch_openml('mnist_784', version=1, return_X_y=True, as_frame=False)
X = X[:1000] / 255.0
y = y[:1000].astype(int)
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2)

# Define a tiny model
model = nn.Sequential(nn.Linear(784, 128), nn.ReLU(), nn.Linear(128, 10))
loss_fn = nn.CrossEntropyLoss()
sim = IonicMLSimulator(d_model=128, learning_rate=0.01)

# Convert data to list of (x,y) tuples
train_data = list(zip(X_train, y_train))

# Training loop (classical, not O(1))
for epoch in range(10):
    for i in range(0, len(train_data), 32):
        batch = train_data[i:i+32]
        loss = sim.train_step(batch, model, loss_fn)
    print(f"Epoch {epoch}, loss {loss:.4f}")