vtava's picture
Upload verified PDelta3-CLVR checkpoint for layers [3, 7, 11]
a38f163 verified
Raw History Blame Contribute Delete
1.68 kB
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