from __future__ import annotations

import argparse
import sys
from pathlib import Path
from typing import List

import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset

# 允许脚本直接运行时进行相对导入
_CUR = Path(__file__).resolve().parent
if str(_CUR) not in sys.path:
    sys.path.insert(0, str(_CUR))
from data import load_parallel_corpus
from preprocess import preprocess_parallel, Vocab, PAD
from models import BiGRUSeqClassifier


def logits_to_tokens(logits: torch.Tensor) -> torch.Tensor:
    return logits.argmax(dim=-1)


def decode_sent(ids: List[int], vocab: Vocab) -> str:
    words = [w for w in vocab.decode(ids) if w != PAD]
    return " ".join(words)


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--epochs", type=int, default=3)
    p.add_argument("--batch-size", type=int, default=128)
    p.add_argument("--hidden", type=int, default=128)
    p.add_argument("--layers", type=int, default=1)
    p.add_argument("--lr", type=float, default=1e-3)
    p.add_argument("--device", type=str, default=None)
    args = p.parse_args()

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

    en_raw, fr_raw = load_parallel_corpus(limit=3000)

    # 预处理：目标长度对齐英文长度
    tmp_en, _, tmp_en_vocab, _ = preprocess_parallel(en_raw, fr_raw)
    ref_len = len(tmp_en[0]) if tmp_en else 0
    en_ids, fr_ids, en_vocab, fr_vocab = preprocess_parallel(en_raw, fr_raw, ref_length=ref_len)

    x = torch.tensor(en_ids, dtype=torch.long)
    y = torch.tensor(fr_ids, dtype=torch.long)

    ds = TensorDataset(x, y)
    dl = DataLoader(ds, batch_size=args.batch_size, shuffle=True)

    model = BiGRUSeqClassifier(
        vocab_size=len(fr_vocab.stoi),
        hidden_size=args.hidden,
        num_layers=args.layers,
        emb_dim=args.hidden,
    ).to(device)

    # 忽略 PAD 的损失
    pad_id = fr_vocab.stoi[PAD]
    criterion = nn.CrossEntropyLoss(ignore_index=pad_id)
    optim = torch.optim.Adam(model.parameters(), lr=args.lr)

    model.train()
    for epoch in range(1, args.epochs + 1):
        total_tok = 0
        loss_sum = 0.0
        for xb, yb in dl:
            xb = xb.to(device)
            yb = yb.to(device)

            logits = model(xb)  # [N, T, V]
            N, T, V = logits.shape
            loss = criterion(logits.view(N * T, V), yb.view(N * T))

            optim.zero_grad()
            loss.backward()
            optim.step()

            loss_sum += float(loss.item()) * (N * T)
            total_tok += int(N * T)

        print(f"[Lab04] Epoch {epoch:02d} | Loss: {loss_sum/max(1,total_tok):.4f}")

    # 输出前 10 条示例
    model.eval()
    out_path = Path("Lab04/result.txt")
    out_path.parent.mkdir(parents=True, exist_ok=True)
    with torch.no_grad(), out_path.open("w", encoding="utf-8") as f:
        take = min(10, x.size(0))
        for i in range(take):
            inp = x[i : i + 1].to(device)
            logits = model(inp)
            pred_ids = logits_to_tokens(logits)[0].tolist()
            src_text = decode_sent(x[i].tolist(), en_vocab)
            tgt_text = decode_sent(pred_ids, fr_vocab)
            f.write(f"{src_text} -> {tgt_text}\n")
    print(f"[Lab04] 示例结果已写入: {out_path}")


if __name__ == "__main__":
    main()
