from __future__ import annotations from pathlib import Path from typing import Any import torch from transformers import AutoTokenizer, PreTrainedTokenizerBase from tiny_router.config import RouterModelConfig from tiny_router.model import TinyRouterModel from tiny_router.runtime import dump_json, load_json def load_tokenizer(model_name_or_path: str | Path) -> PreTrainedTokenizerBase: return AutoTokenizer.from_pretrained(model_name_or_path, use_fast=False) def save_checkpoint( output_dir: str | Path, model: TinyRouterModel, tokenizer, model_config: RouterModelConfig, training_args: dict[str, Any], metrics: dict[str, Any] | None = None, ) -> None: output_dir = Path(output_dir) output_dir.mkdir(parents=True, exist_ok=True) torch.save(model.state_dict(), output_dir / "model.pt") tokenizer.save_pretrained(output_dir) dump_json(output_dir / "model_config.json", model_config.to_dict()) dump_json(output_dir / "training_args.json", training_args) if metrics is not None: dump_json(output_dir / "metrics.json", metrics) def load_checkpoint(model_dir: str | Path, device: torch.device) -> tuple[TinyRouterModel, Any, RouterModelConfig]: model_dir = Path(model_dir) model_config = RouterModelConfig.from_dict(load_json(model_dir / "model_config.json")) tokenizer = load_tokenizer(model_dir) model = TinyRouterModel(model_config) state_dict = torch.load(model_dir / "model.pt", map_location=device, weights_only=True) model.load_state_dict(state_dict) model.to(device) model.eval() return model, tokenizer, model_config def save_temperature_scaling( output_dir: str | Path, temperature_scaling: dict[str, Any], ) -> None: dump_json(Path(output_dir) / "temperature_scaling.json", temperature_scaling) def load_temperature_scaling(model_dir: str | Path) -> dict[str, float]: path = Path(model_dir) / "temperature_scaling.json" if not path.exists(): return {} payload = load_json(path) return {head: float(value) for head, value in payload.get("per_head", {}).items()}