File size: 1,682 Bytes
a38f163
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
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