import torch
from transformers import AutoTokenizer, AutoModelForCausalLM
import numpy as np
from scipy.io.wavfile import read

# 1. Initialize SmolLM2-135M on CPU/GPU for the meta-optimizer
model_id = "HuggingFaceTB/SmolLM2-135M-Instruct"
device = "cuda" if torch.cuda.is_available() else "cpu"

print(f"Loading {model_id} as the meta-optimizer on {device}...")
tokenizer = AutoTokenizer.from_pretrained(model_id)
llm_model = AutoModelForCausalLM.from_pretrained(model_id).to(device)

# Map single characters to our target optimizer steps
optimizer_actions = {"A": "sgd", "B": "rk4"}
token_map = {letter: tokenizer.convert_tokens_to_ids(letter) for letter in optimizer_actions.keys()}
token_ids = list(token_map.values())
id_to_letter = {v: k for k, v in token_map.items()}

# --- Mock Data Setup (Using random data here as a placeholder for raw audio tensors) ---

from scipy.io.wavfile import read

# Load training and test data
X_train = read('../X_train.wav')[1].reshape(-1, 784)
y_train = (read('../y_train.wav')[1] * 9).astype(int)
X_test = read('../X_test.wav')[1].reshape(-1, 784)
y_test = (read('../y_test.wav')[1] * 9).astype(int)
X20 = X_test[:1000]
yt20 = y_test[:1000]




# --- Your MLP Classifier Architecture ---
class MLPClassifier:
    def __init__(self, input_size, hidden_size, output_size, learning_rate=None):
        self.W1 = np.random.randn(input_size, hidden_size) * 0.01
        self.b1 = np.zeros((1, hidden_size))
        self.W2 = np.random.randn(hidden_size, output_size) * 0.01
        self.b2 = np.zeros((1, output_size))
        self.learning_rate = learning_rate if learning_rate is not None else np.random.rand(4)
    
    def relu(self, x): return np.maximum(0, x)
    def relu_derivative(self, x): return np.where(x > 0, 1, 0)
    def softmax(self, x):
        exp_x = np.exp(x - np.max(x, axis=1, keepdims=True))
        return exp_x / np.sum(exp_x, axis=1, keepdims=True)
    
    def _forward_with_weights(self, X, W1, b1, W2, b2):
        z1 = np.dot(X, W1) + b1
        a1 = self.relu(z1)
        z2 = np.dot(a1, W2) + b2
        return self.softmax(z2), z1, a1
    
    def _backward_with_weights(self, X, y_true, y_pred, z1, a1, W2):
        m = y_true.shape[0]
        dz2 = y_pred - y_true
        dW2 = np.dot(a1.T, dz2) / m
        db2 = np.sum(dz2, axis=0, keepdims=True) / m
        da1 = np.dot(dz2, W2.T)
        dz1 = da1 * self.relu_derivative(z1)
        dW1 = np.dot(X.T, dz1) / m
        db1 = np.sum(dz1, axis=0, keepdims=True) / m
        return dW1, db1, dW2, db2

    def compute_gradients(self, X, y_true, W1, b1, W2, b2):
        y_pred, z1, a1 = self._forward_with_weights(X, W1, b1, W2, b2)
        return self._backward_with_weights(X, y_true, y_pred, z1, a1, W2)
    
    def update_sgd(self, X, y_true):
        dW1, db1, dW2, db2 = self.compute_gradients(X, y_true, self.W1, self.b1, self.W2, self.b2)
        self.W1 -= self.learning_rate[0] * dW1
        self.b1 -= self.learning_rate[1] * db1
        self.W2 -= self.learning_rate[2] * dW2
        self.b2 -= self.learning_rate[3] * db2

    def update_rk4(self, X, y_true):
        lr = self.learning_rate
        dW1_k1, db1_k1, dW2_k1, db2_k1 = self.compute_gradients(X, y_true, self.W1, self.b1, self.W2, self.b2)
        dW1_k2, db1_k2, dW2_k2, db2_k2 = self.compute_gradients(X, y_true, self.W1 - 0.5 * lr[0] * dW1_k1, self.b1 - 0.5 * lr[1] * db1_k1, self.W2 - 0.5 * lr[2] * dW2_k1, self.b2 - 0.5 * lr[3] * db2_k1)
        dW1_k3, db1_k3, dW2_k3, db2_k3 = self.compute_gradients(X, y_true, self.W1 - 0.5 * lr[0] * dW1_k2, self.b1 - 0.5 * lr[1] * db1_k2, self.W2 - 0.5 * lr[2] * dW2_k2, self.b2 - 0.5 * lr[3] * db2_k2)
        dW1_k4, db1_k4, dW2_k4, db2_k4 = self.compute_gradients(X, y_true, self.W1 - lr[0] * dW1_k3, self.b1 - lr[1] * db1_k3, self.W2 - lr[2] * dW2_k3, self.b2 - lr[3] * db2_k3)
        
        self.W1 -= (lr[0] / 6.0) * (dW1_k1 + 2 * dW1_k2 + 2 * dW1_k3 + dW1_k4)
        self.b1 -= (lr[1] / 6.0) * (db1_k1 + 2 * db1_k2 + 2 * db1_k3 + db1_k4)
        self.W2 -= (lr[2] / 6.0) * (dW2_k1 + 2 * dW2_k2 + 2 * dW2_k3 + dW2_k4)
        self.b2 -= (lr[3] / 6.0) * (db2_k1 + 2 * db2_k2 + 2 * db2_k3 + db2_k4)

    def predict(self, X):
        probabilities, _, _ = self._forward_with_weights(X, self.W1, self.b1, self.W2, self.b2)
        return np.argmax(probabilities, axis=1)

    def score(self, X, y_true):
        return np.mean(self.predict(X) == y_true)


# --- Meta-Learning Loop ---
if __name__ == "__main__":
    learning_rate = np.array([0.05, 0.05, 0.05, 0.05])
    mlp = MLPClassifier(input_size=784, hidden_size=100, output_size=10, learning_rate=learning_rate)

    print("Beginning LLM-guided optimizer tracking...\n")
    
    current_score = 0.0
    last_strategy = "None"
    
    for step in range(1, 100000):
        idx = np.random.randint(0, 60000, 100)
        X = X_train[idx]
        yt = y_train[idx]
        
        current_score = mlp.score(X, yt)
        print(f"--- Global Step {step} ---")
        print(f"Current Batch Training Accuracy: {current_score:.4f}")
        print(f"Last Strategy Picked: {last_strategy}")

        # Construct a prompt passing numerical state representations to SmolLM2
        prompt = (
            f"<|im_start|>system\n"
            f"You are choosing neural network optimization moves.\n"
            f"Current classification performance metric is {current_score:.4f}.\n"
            f"Select A for basic Stochastic Gradient Descent (SGD).\n"
            f"Select B for higher order Runge-Kutta 4 (RK4).\n"
            f"Which execution type is ideal? Output letter:<|im_end|>\n"
            f"<|im_start|>assistant\n"
        )
        
        inputs = tokenizer(prompt, return_tensors="pt").to(device)
        
        with torch.no_grad():
            outputs = llm_model(**inputs)
            logits = outputs.logits[0, -1, :]
            
            # Slice only our targets ('A' vs 'B')
            target_logits = logits[token_ids]
            probs = torch.softmax(target_logits, dim=-1).float().cpu().numpy()
            
            chosen_token = token_ids[np.argmax(probs)]
            chosen_letter = id_to_letter[chosen_token]
            strategy = optimizer_actions[chosen_letter]

        print(f"SmolLM2 Meta Logic Probabilities -> A(SGD): {probs[0]:.2f} | B(RK4): {probs[1]:.2f}")
        print(f"Action Taken: Applying {strategy.upper()} update step.")
        
        # Execute the optimization step chosen by the model's logits
        if strategy == "sgd":
            mlp.update_sgd(X, np.eye(10)[yt])
        elif strategy == "rk4":
            mlp.update_rk4(X, np.eye(10)[yt])
            
        last_strategy = strategy.upper()
        print("\n")
