import argparse
import json
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))

import torch
from PIL import Image

from common.classification_training import build_model, choose_device, default_transforms


def softmax_probabilities(logits: torch.Tensor, labels: dict[int, str]) -> dict[str, float]:
    probs = torch.softmax(logits, dim=1)[0].tolist()
    return {labels[idx]: float(prob) for idx, prob in enumerate(probs)}


def predict_density(image_path: str | Path, model_path: str | Path, device_name: str = "auto") -> dict:
    image_id = Path(image_path).stem.replace("_potcrop", "")
    checkpoint = torch.load(model_path, map_location="cpu")
    labels = {int(key): value for key, value in checkpoint["labels"].items()}
    image_size = int(checkpoint.get("image_size", 224))
    model = build_model(checkpoint["model_name"], int(checkpoint["num_classes"]))
    model.load_state_dict(checkpoint["model_state_dict"])
    model.eval()
    device = choose_device(device_name)
    model.to(device)
    image = Image.open(image_path).convert("RGB")
    tensor = default_transforms(image_size)(image).unsqueeze(0).to(device)
    with torch.inference_mode():
        logits = model(tensor)
    label_id = int(logits.argmax(dim=1).item())
    probabilities = softmax_probabilities(logits.cpu(), labels)
    confidence = probabilities[labels[label_id]]
    return {
        "image_id": image_id,
        "predicted_label": labels[label_id],
        "predicted_label_id": label_id,
        "confidence": confidence,
        "probabilities": probabilities,
    }


def main() -> int:
    parser = argparse.ArgumentParser(description="Predict soup density label from pot crop image and steam level.")
    parser.add_argument("--image", required=True)
    parser.add_argument("--model-path", default="models/density/density_classifier_extracted_v1/model.pt")
    parser.add_argument("--device", default="auto")
    parser.add_argument("--output", default=None)
    args = parser.parse_args()
    result = predict_density(args.image, args.model_path, args.device)
    text = json.dumps(result, ensure_ascii=False, indent=2)
    if args.output:
        Path(args.output).parent.mkdir(parents=True, exist_ok=True)
        Path(args.output).write_text(text, encoding="utf-8")
    print(text)
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
