feat(llm): Qwen QLoRA 4-bit 真微调路径 + RAG 编码冻结/微调分流

- use_lora 时 Qwen 改 4-bit(NF4) 加载,配 prepare_model_for_kbit_training 使 LoRA 可回传梯度
- _encode_rag 拆分:冻结走 no_grad+缓存,微调走保留计算图+不缓存
- fold_0/1 delivery R2 0.44/0.24,较冻结版 0.185 显著提升
This commit is contained in:
Michelle0474 2026-07-01 15:52:20 +08:00
parent 0fc72a9474
commit 7248148515

View File

@ -54,8 +54,19 @@ class LLMPromptEncoder(nn.Module):
self.encoder = T5EncoderModel.from_pretrained(model_name_or_path) self.encoder = T5EncoderModel.from_pretrained(model_name_or_path)
self.hidden_size = self.encoder.config.d_model self.hidden_size = self.encoder.config.d_model
elif _is_qwen: elif _is_qwen:
self.encoder = AutoModel.from_pretrained( # 微调(use_lora)时用 4-bit 量化(QLoRA)省显存;纯冻结时用 fp16
model_name_or_path, trust_remote_code=True, torch_dtype=torch.float16) if use_lora:
from transformers import BitsAndBytesConfig
_bnb = BitsAndBytesConfig(
load_in_4bit=True, bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.float16,
bnb_4bit_use_double_quant=True)
self.encoder = AutoModel.from_pretrained(
model_name_or_path, trust_remote_code=True,
quantization_config=_bnb, device_map={"": 0})
else:
self.encoder = AutoModel.from_pretrained(
model_name_or_path, trust_remote_code=True, torch_dtype=torch.float16)
self.hidden_size = self.encoder.config.hidden_size self.hidden_size = self.encoder.config.hidden_size
else: else:
self.encoder = AutoModel.from_pretrained(model_name_or_path) self.encoder = AutoModel.from_pretrained(model_name_or_path)
@ -82,6 +93,10 @@ class LLMPromptEncoder(nn.Module):
def _apply_lora(self, r: int, alpha: int, dropout: float) -> None: def _apply_lora(self, r: int, alpha: int, dropout: float) -> None:
from peft import LoraConfig, get_peft_model from peft import LoraConfig, get_peft_model
# 4-bit 量化模型(QLoRA)需先 prepare才能正确接收梯度
if getattr(self.encoder, "is_loaded_in_4bit", False) or getattr(self.encoder, "is_loaded_in_8bit", False):
from peft import prepare_model_for_kbit_training
self.encoder = prepare_model_for_kbit_training(self.encoder)
for p in self.encoder.parameters(): for p in self.encoder.parameters():
p.requires_grad = False p.requires_grad = False
# 不同架构注意力层命名不同: # 不同架构注意力层命名不同:
@ -153,34 +168,39 @@ class LLMPromptEncoder(nn.Module):
"into your internal representation." "into your internal representation."
) )
@torch.no_grad() def _rag_encode_batch(self, prompts, device):
"""编码一批 RAG prompt取最后有效 token。grad 由调用方上下文决定。"""
outs = []
for i in range(0, len(prompts), 4):
bt = prompts[i:i+4]
enc = self.tokenizer(
bt, padding=True, truncation=True,
max_length=self.max_length, return_tensors="pt",
).to(device)
out = self.encoder(**enc).last_hidden_state # [B,L,H]
lengths = enc["attention_mask"].sum(1) - 1
b = torch.arange(out.size(0), device=device)
outs.append(out[b, lengths.long(), :].float()) # [B,H]
return torch.cat(outs, 0)
def _encode_rag(self, smiles, device): def _encode_rag(self, smiles, device):
"""RAG 编码每个分子→检索→构造prompt→Qwen编码→取最后有效token。带缓存。""" """RAG 编码。冻结时no_grad + 缓存(快)。微调时:算梯度 + 不缓存。"""
feats = [] if self._frozen:
to_compute = [] # 冻结路径缓存复用no_grad
keys = [] keys = [f"RAG::{self._rag_pool_id}::{s}" for s in smiles]
for s in smiles: to_compute = [s for s, k in zip(smiles, keys) if k not in self._cache]
key = f"RAG::{self._rag_pool_id}::{s}" if to_compute:
keys.append(key) uniq = list(dict.fromkeys(to_compute))
if key not in self._cache: prompts = [self._build_rag_prompt(s, self._retrieve_topk(s)) for s in uniq]
to_compute.append(s) with torch.no_grad():
# 逐个构造 prompt检索是 CPU 操作) feat = self._rag_encode_batch(prompts, device)
if to_compute: for s, v in zip(uniq, feat):
uniq = list(dict.fromkeys(to_compute))
prompts = [self._build_rag_prompt(s, self._retrieve_topk(s)) for s in uniq]
for i in range(0, len(uniq), 4):
bt = prompts[i:i+4]
enc = self.tokenizer(
bt, padding=True, truncation=True,
max_length=self.max_length, return_tensors="pt",
).to(device)
out = self.encoder(**enc).last_hidden_state # [B,L,H]
lengths = enc["attention_mask"].sum(1) - 1 # 最后有效token位置
b = torch.arange(out.size(0), device=device)
last_tok = out[b, lengths.long(), :] # [B,H]
for s, v in zip(uniq[i:i+4], last_tok):
self._cache[f"RAG::{self._rag_pool_id}::{s}"] = v.float().cpu() self._cache[f"RAG::{self._rag_pool_id}::{s}"] = v.float().cpu()
return torch.stack([self._cache[k] for k in keys]).to(device) return torch.stack([self._cache[k] for k in keys]).to(device)
else:
# 微调路径:每次重新编码,保留计算图(算梯度),不缓存
prompts = [self._build_rag_prompt(s, self._retrieve_topk(s)) for s in smiles]
return self._rag_encode_batch(prompts, device)
def _mean_pool(self, last_hidden: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: def _mean_pool(self, last_hidden: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
m = mask.unsqueeze(-1).float() # [B, L, 1] m = mask.unsqueeze(-1).float() # [B, L, 1]