import argparse
import json
import sys
from pathlib import Path

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

import yaml

from crop.rfdetr_adapter import load_rfdetr_seg_model


if hasattr(sys.stdout, "reconfigure"):
    sys.stdout.reconfigure(encoding="utf-8", errors="replace")
if hasattr(sys.stderr, "reconfigure"):
    sys.stderr.reconfigure(encoding="utf-8", errors="replace")


def disable_lightning_checkpoints() -> None:
    import rfdetr.training as rfdetr_training
    from pytorch_lightning.callbacks import ModelCheckpoint

    original_build_trainer = rfdetr_training.build_trainer

    def build_trainer_without_lightning_checkpoints(*args, **kwargs):
        trainer = original_build_trainer(*args, **kwargs)
        trainer.callbacks = [callback for callback in trainer.callbacks if not isinstance(callback, ModelCheckpoint)]
        return trainer

    rfdetr_training.build_trainer = build_trainer_without_lightning_checkpoints


def main() -> int:
    parser = argparse.ArgumentParser(description="Train RF-DETR-Seg on COCO segmentation pot_inner dataset.")
    parser.add_argument("--config", default="configs/crop_rfdetr_seg.yaml")
    parser.add_argument("--dataset-dir", default=None)
    parser.add_argument("--epochs", type=int, default=None)
    parser.add_argument("--batch-size", type=int, default=None)
    parser.add_argument("--grad-accum-steps", type=int, default=None)
    parser.add_argument("--lr", type=float, default=None)
    parser.add_argument("--output-dir", default=None)
    parser.add_argument("--report-dir", default=None)
    parser.add_argument("--device", default=None)
    parser.add_argument("--resume", default=None)
    parser.add_argument("--disable-lightning-checkpoints", action="store_true")
    args = parser.parse_args()

    with Path(args.config).open("r", encoding="utf-8") as f:
        config = yaml.safe_load(f)
    dataset_dir = args.dataset_dir or config["dataset_dir"]
    model_size = config.get("model_size", "small")
    if model_size not in config.get("allowed_model_sizes", ["small", "medium"]):
        raise ValueError("Only RF-DETR-Seg Small/Medium are enabled until license review is complete.")

    output_dir = Path(args.output_dir or config["output_dir"])
    report_dir = Path(args.report_dir or config["report_dir"])
    output_dir.mkdir(parents=True, exist_ok=True)
    report_dir.mkdir(parents=True, exist_ok=True)

    model = load_rfdetr_seg_model(model_size=model_size, checkpoint_path=config.get("checkpoint_path"))
    train_kwargs = {
        "dataset_dir": dataset_dir,
        "epochs": int(args.epochs if args.epochs is not None else config["epochs"]),
        "batch_size": int(args.batch_size if args.batch_size is not None else config["batch_size"]),
        "grad_accum_steps": int(args.grad_accum_steps if args.grad_accum_steps is not None else config.get("grad_accum_steps", 1)),
        "lr": float(args.lr if args.lr is not None else config["lr"]),
        "output_dir": str(output_dir),
        "tensorboard": bool(config.get("tensorboard", False)),
        "wandb": bool(config.get("wandb", False)),
    }
    resume = args.resume if args.resume is not None else config.get("resume")
    if resume:
        train_kwargs["resume"] = resume
    device = args.device if args.device is not None else config.get("device")
    if device and device != "auto":
        train_kwargs["device"] = device

    if args.disable_lightning_checkpoints or config.get("disable_lightning_checkpoints", False):
        disable_lightning_checkpoints()

    (report_dir / "train_config.json").write_text(json.dumps(train_kwargs, indent=2), encoding="utf-8")
    model.train(**train_kwargs)
    print(f"training finished. output_dir={output_dir}")
    return 0


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