File size: 7,155 Bytes
09ccad2 | 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 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 | """Qwen3.5-4B 学生模型加载 + 隐藏态提取 + ComfyUI 感知的显存 offload。
load_student 移植自训练管线 extract_student_features.py:
- 修复 checkpoint 的 `model.` 多余前缀(否则全层随机初始化)
- 隐藏态取 hidden_states[-1](tie_last_hidden_states -> post-final-norm,与训练/评估同源)
"""
from __future__ import annotations
import json
import os
from collections import defaultdict
import torch
try: # ComfyUI 内运行
import comfy.model_management as comfy_mm
def get_torch_device() -> torch.device:
return comfy_mm.get_torch_device()
def soft_empty_cache() -> None:
comfy_mm.soft_empty_cache()
except ImportError: # 独立测试环境
def get_torch_device() -> torch.device:
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
def soft_empty_cache() -> None:
if torch.cuda.is_available():
torch.cuda.empty_cache()
def load_student(model_dir: str) -> torch.nn.Module:
"""加载 Qwen3.5-4B 权重(CPU),剥离 `model.` 前缀。
model_dir 为 .gguf 文件时走自研加载器(gguf_qwen35):
读取 + 反量化 + 张量名映射均自实现,不依赖 transformers 的 GGUF 支持,
产出与训练 checkpoint 同格式的 state dict。
"""
from safetensors import safe_open
from transformers import AutoConfig, AutoModel
if model_dir.endswith(".gguf"):
try:
from . import gguf_qwen35
except ImportError:
import gguf_qwen35
# 生产路径: 量化块驻留 + 前向逐层反量化(显存 ≈ 量化大小而非 fp16 展开)
model, info = gguf_qwen35.load_gguf_model_quantized(model_dir)
print(f" [student] GGUF loaded: {info}", flush=True)
return model
cfg = AutoConfig.from_pretrained(model_dir)
model = AutoModel.from_config(cfg) # 无权重实例化,避免 from_pretrained 错误初始化
index = json.load(open(os.path.join(model_dir, "model.safetensors.index.json")))
by_shard: dict[str, list[str]] = defaultdict(list)
for k, sh in index["weight_map"].items():
by_shard[sh].append(k)
sd: dict[str, torch.Tensor] = {}
for shard, keys in by_shard.items():
with safe_open(os.path.join(model_dir, shard), framework="pt", device="cpu") as sf:
for k in keys:
newk = k[len("model."):] if k.startswith("model.") else k
sd[newk] = sf.get_tensor(k)
missing, unexpected = model.load_state_dict(sd, strict=False)
del sd
assert not missing, f"missing after prefix strip: {missing[:10]}"
if unexpected:
print(f" [student] ignored {len(unexpected)} unexpected keys (e.g. {unexpected[:3]})", flush=True)
return model
def find_language_model(model: torch.nn.Module) -> torch.nn.Module:
for name in ("language_model", "model", "text_model"):
if hasattr(model, name):
sub = getattr(model, name)
if hasattr(sub, "layers") or hasattr(sub, "config"):
return sub
raise RuntimeError(f"cannot locate language model submodule; attrs={[n for n in dir(model) if not n.startswith('_')]}")
def _move_plain_tensors(module: torch.nn.Module, device: torch.device) -> None:
"""迁移普通 tensor 属性(GGMLTensor 量化权重/embed 权重等非 Parameter/buffer)。
nn.Module.to() 只迁移 Parameter 与 registered buffer,量化权重是普通属性,
必须手动搬——否则 encode 后 offload 不彻底,量化权重滞留 GPU。
"""
for name, attr in list(module.__dict__.items()):
if isinstance(attr, torch.Tensor) and not isinstance(attr, torch.nn.Parameter):
setattr(module, name, attr.to(device))
for child in module.children():
_move_plain_tensors(child, device)
class StudentTextEncoder:
"""学生模型封装: prompt -> [S_S, 2560] bf16 hidden(post-final-norm)。
gpu_mem=""(默认): 整模型进 GPU;encode 后 lowvram 搬回 CPU。
gpu_mem="5GiB"/"16GiB" 等: accelerate 层间 offload(权重驻留 RAM、按层流式进 GPU),
适合 24GB 及以下卡(33B DiT + 4B 学生错峰)。
lowvram=True(默认): encode 完成后立即释放 GPU 占用,让位给 DiT 采样。
"""
def __init__(self, model_dir: str, dtype: torch.dtype = torch.bfloat16,
lowvram: bool = True, gpu_mem: str = ""):
self.model_dir = model_dir
self.dtype = dtype
self.lowvram = lowvram
self.gpu_mem = gpu_mem
self._dispatched = False
self._model: torch.nn.Module | None = None
self._lm: torch.nn.Module | None = None
self._tok = None
def _ensure_ready(self):
if self._model is None:
from transformers import AutoTokenizer
model = load_student(self.model_dir).to(self.dtype)
if self.gpu_mem:
from accelerate import dispatch_model, infer_auto_device_map
max_memory = {"cpu": "20GiB", 0: self.gpu_mem}
no_split = getattr(model, "_no_split_modules", None) or None
device_map = infer_auto_device_map(model, max_memory=max_memory,
no_split_module_classes=no_split)
dispatch_model(model, device_map=device_map)
self._dispatched = True
self._model = model
self._lm = find_language_model(self._model)
if self.model_dir.endswith(".gguf"):
tok_dir = os.path.dirname(self.model_dir)
if not os.path.exists(os.path.join(tok_dir, "tokenizer.json")):
raise RuntimeError(
f"GGUF 同目录 {tok_dir} 缺少 tokenizer.json —— 请把 Qwen3.5-4B 的 tokenizer 文件"
"(tokenizer.json / tokenizer_config.json / vocab.json / merges.txt / chat_template.jinja)拷到该目录")
self._tok = AutoTokenizer.from_pretrained(tok_dir)
else:
self._tok = AutoTokenizer.from_pretrained(self.model_dir)
print(f" [student] loaded {self.model_dir} (gpu_mem={self.gpu_mem or 'whole'})", flush=True)
dev = get_torch_device()
if not self._dispatched and next(self._lm.parameters()).device != dev:
self._model.to(dev)
def offload(self) -> None:
if self._lm is None:
return
if self._dispatched:
soft_empty_cache()
return
self._model.to("cpu")
_move_plain_tensors(self._model, torch.device("cpu"))
soft_empty_cache()
@torch.no_grad()
def __call__(self, prompt: str) -> torch.Tensor:
self._ensure_ready()
dev = get_torch_device()
ids = self._tok(prompt, add_special_tokens=False)["input_ids"]
out = self._lm(input_ids=torch.tensor([ids], device=dev), output_hidden_states=True)
h = out.hidden_states[-1][0].to(self.dtype) # [S_S, 2560] post-final-norm
if self.lowvram:
self.offload()
return h
|