#!/usr/bin/env python3
"""Samodzielny skrypt inferencyjny dla pakietu MicroPLLM.

Użycie:
    python run_inference.py "Polska to kraj"
    python run_inference.py --prompt "Czym jest sztuczna inteligencja?" --max_tokens 100 --temp 0.7
"""

import argparse
import os
import sys
import time
import torch
import torch.nn.functional as F
import sentencepiece as spm
from model import MicroPLLM


def main():
    parser = argparse.ArgumentParser(description="Testowa inferencja MicroPLLM (~109M)")
    parser.add_argument("prompt_pos", nargs="?", default=None, help="Prompt tekstowy (pozycyjny)")
    parser.add_argument("--prompt", "-p", default="Polska to kraj", help="Prompt tekstowy")
    parser.add_argument("--weights", "-w", default="micro_pl_lm_weights.pt", help="Ścieżka do wag .pt")
    parser.add_argument("--tokenizer", "-t", default="pl_bpe_32k.model", help="Ścieżka do tokenizera")
    parser.add_argument("--max_tokens", "-n", type=int, default=100, help="Maksymalna liczba nowych tokenów")
    parser.add_argument("--temp", type=float, default=0.7, help="Temperatura samplowania")
    parser.add_argument("--top_p", type=float, default=0.9, help="Top-p sampling")
    parser.add_argument("--top_k", type=int, default=50, help="Top-k sampling")
    parser.add_argument("--rep_penalty", type=float, default=1.15, help="Repetition penalty")
    parser.add_argument("--device", default=None, help="Urządzenie (cuda / cpu / mps)")

    args = parser.parse_args()
    prompt = args.prompt_pos if args.prompt_pos else args.prompt

    device = args.device
    if device is None:
        device = "cuda" if torch.cuda.is_available() else "cpu"

    # Inteligentne wykrywanie ścieżki tokenizera
    tok_path = args.tokenizer
    if not os.path.exists(tok_path):
        for candidate in ["tokenizer/pl_bpe_32k.model", "/media/disk5/micro_pl_lm/tokenizer/pl_bpe_32k.model", "/media/fujitsuserv/www/pl_bpe_32k.model"]:
            if os.path.exists(candidate):
                tok_path = candidate
                break

    print(f"-> Ładowanie tokenizera z {tok_path}...")
    sp = spm.SentencePieceProcessor()
    sp.load(tok_path)

    # Inteligentne wykrywanie ścieżki wag
    w_path = args.weights
    if not os.path.exists(w_path):
        for candidate in ["/media/fujitsuserv/www/micro_pl_lm_weights.pt", "checkpoints/latest.pt", "/media/disk5/micro_pl_lm/checkpoints/latest.pt"]:
            if os.path.exists(candidate):
                w_path = candidate
                break

    print(f"-> Inicjalizacja modelu MicroPLLM na {device} z wag {w_path}...")
    model = MicroPLLM()
    state_dict = torch.load(w_path, map_location=device, weights_only=False)
    if "model" in state_dict:
        state_dict = state_dict["model"]
    model.load_state_dict(state_dict, strict=True)
    model.to(device)
    model.eval()

    print(f"-> Generowanie tekstu dla promptu: \"{prompt}\"")
    toks = sp.encode(prompt, out_type=int)
    if not toks:
        toks = [sp.bos_id() if sp.bos_id() >= 0 else 1]

    idx = torch.tensor([toks], dtype=torch.long, device=device)
    t0 = time.time()
    generated = []

    with torch.no_grad():
        for _ in range(args.max_tokens):
            idx_cond = idx if idx.size(1) <= 2048 else idx[:, -2048:]
            logits = model(idx_cond)
            next_logit = logits[:, -1, :].clone().float()

            if args.rep_penalty != 1.0:
                for tid in set(idx[0].tolist()):
                    if next_logit[0, tid] < 0:
                        next_logit[0, tid] *= args.rep_penalty
                    else:
                        next_logit[0, tid] /= args.rep_penalty

            if args.temp <= 0.05:
                next_tok = torch.argmax(next_logit, dim=-1, keepdim=True)
            else:
                next_logit = next_logit / max(1e-4, args.temp)
                if args.top_k > 0:
                    v, _ = torch.topk(next_logit, min(args.top_k, next_logit.size(-1)))
                    next_logit[next_logit < v[:, [-1]]] = -float("Inf")
                probs = F.softmax(next_logit, dim=-1)
                next_tok = torch.multinomial(probs, num_samples=1)

            token_id = next_tok.item()
            if token_id == sp.eos_id():
                break

            generated.append(token_id)
            idx = torch.cat((idx, next_tok), dim=1)

    elapsed = time.time() - t0
    full_text = sp.decode(idx[0].tolist())
    gen_text = sp.decode(generated)

    print("\n" + "=" * 50)
    print("WYGENEROWANY TEKST:")
    print("=" * 50)
    print(full_text)
    print("=" * 50)
    print(f"Metryki: {len(generated)} tokenów w {elapsed:.2f}s ({len(generated)/max(1e-3, elapsed):.1f} tok/s)\n")


if __name__ == "__main__":
    main()
