from __future__ import annotations import json import sys from pathlib import Path import torch from safetensors.torch import load_file from transformers import AutoTokenizer, Qwen3_5ForCausalLM def load_model(path=None, device="cpu", dtype=None): root = Path(path or Path(__file__).resolve().parent) for p in (root / "src", root, root / "scripts"): if str(p) not in sys.path: sys.path.insert(0, str(p)) from train_qwen35_pdelta3_clvr_sequential import QwenPDelta3CLVRConfig, replace_full_attention_layers meta = json.loads((root / "tinycenn_qwen35.json").read_text()) accepted = [int(x) for x in meta["accepted_full_attention_layers"]] cfg = QwenPDelta3CLVRConfig.from_dict(meta["replacement_config"]) if dtype is None: dtype = torch.bfloat16 if device.startswith("cuda") and torch.cuda.is_bf16_supported() else (torch.float16 if device.startswith("cuda") else torch.float32) model = Qwen3_5ForCausalLM.from_pretrained(root, dtype=dtype, local_files_only=True, attn_implementation="eager") replace_full_attention_layers(model, cfg, accepted) single = root / "model.safetensors" if single.exists(): model.load_state_dict(load_file(str(single), device="cpu"), strict=False) else: index = json.loads((root / "model.safetensors.index.json").read_text()) for shard in sorted(set(index["weight_map"].values())): model.load_state_dict(load_file(str(root / shard), device="cpu"), strict=False) model.config.use_cache = False model.to(device).eval() tokenizer = AutoTokenizer.from_pretrained(root, local_files_only=True, use_fast=True) return model, tokenizer