#!/usr/bin/env python
# predict_nn.py
import sys
import numpy as np
import torch
from train_nn import SimpleNet
from train import build_features, load_csv

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

SCORE_MULTIPLIERS = {5: 1.5, 6: 2.5, 0: 1.0, 3: 5.0}


def main():
    model_file, input_csv, output_csv = sys.argv[1:]

    print("Loading model...")
    checkpoint = torch.load(model_file, map_location=device, weights_only=False)
    model = SimpleNet(checkpoint['in_dim']).to(device)
    model.load_state_dict(checkpoint['state_dict'])
    model.eval()

    mean              = checkpoint['mean']
    std               = checkpoint['std']
    eng_start         = checkpoint['eng_features_start']

    print("Loading data...")
    labels, X = load_csv(input_csv)

    print("Building features...")
    Xf_full = build_features(X)
    Xf      = Xf_full[:, eng_start:]
    Xf_norm = ((Xf - mean) / std).astype(np.float32)

    print("Predicting...")
    preds = []
    with torch.no_grad():
        for i in range(0, len(Xf_norm), 512):
            xb    = torch.tensor(Xf_norm[i:i+512]).to(device)
            probs = torch.softmax(model(xb), dim=1).cpu().numpy()
            for p in probs:
                for cls, mult in SCORE_MULTIPLIERS.items():
                    p[cls] *= mult
                preds.append(int(np.argmax(p)))

    assert len(preds) == len(labels), f"Length mismatch: {len(preds)} vs {len(labels)}"

    print("Writing output...")
    with open(output_csv, 'w') as f:
        f.write("label\n")
        for p in preds:
            f.write(f"{p}\n")
    print("Done.")


if __name__ == "__main__":
    main()