mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-07-22 05:49:55 +08:00
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:
parent
0fc72a9474
commit
7248148515
@ -54,6 +54,17 @@ 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:
|
||||||
|
# 微调(use_lora)时用 4-bit 量化(QLoRA)省显存;纯冻结时用 fp16
|
||||||
|
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(
|
self.encoder = AutoModel.from_pretrained(
|
||||||
model_name_or_path, trust_remote_code=True, torch_dtype=torch.float16)
|
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
|
||||||
@ -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):
|
||||||
def _encode_rag(self, smiles, device):
|
"""编码一批 RAG prompt,取最后有效 token。grad 由调用方上下文决定。"""
|
||||||
"""RAG 编码:每个分子→检索→构造prompt→Qwen编码→取最后有效token。带缓存。"""
|
outs = []
|
||||||
feats = []
|
for i in range(0, len(prompts), 4):
|
||||||
to_compute = []
|
|
||||||
keys = []
|
|
||||||
for s in smiles:
|
|
||||||
key = f"RAG::{self._rag_pool_id}::{s}"
|
|
||||||
keys.append(key)
|
|
||||||
if key not in self._cache:
|
|
||||||
to_compute.append(s)
|
|
||||||
# 逐个构造 prompt(检索是 CPU 操作)
|
|
||||||
if to_compute:
|
|
||||||
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]
|
bt = prompts[i:i+4]
|
||||||
enc = self.tokenizer(
|
enc = self.tokenizer(
|
||||||
bt, padding=True, truncation=True,
|
bt, padding=True, truncation=True,
|
||||||
max_length=self.max_length, return_tensors="pt",
|
max_length=self.max_length, return_tensors="pt",
|
||||||
).to(device)
|
).to(device)
|
||||||
out = self.encoder(**enc).last_hidden_state # [B,L,H]
|
out = self.encoder(**enc).last_hidden_state # [B,L,H]
|
||||||
lengths = enc["attention_mask"].sum(1) - 1 # 最后有效token位置
|
lengths = enc["attention_mask"].sum(1) - 1
|
||||||
b = torch.arange(out.size(0), device=device)
|
b = torch.arange(out.size(0), device=device)
|
||||||
last_tok = out[b, lengths.long(), :] # [B,H]
|
outs.append(out[b, lengths.long(), :].float()) # [B,H]
|
||||||
for s, v in zip(uniq[i:i+4], last_tok):
|
return torch.cat(outs, 0)
|
||||||
|
|
||||||
|
def _encode_rag(self, smiles, device):
|
||||||
|
"""RAG 编码。冻结时:no_grad + 缓存(快)。微调时:算梯度 + 不缓存。"""
|
||||||
|
if self._frozen:
|
||||||
|
# 冻结路径:缓存复用,no_grad
|
||||||
|
keys = [f"RAG::{self._rag_pool_id}::{s}" for s in smiles]
|
||||||
|
to_compute = [s for s, k in zip(smiles, keys) if k not in self._cache]
|
||||||
|
if to_compute:
|
||||||
|
uniq = list(dict.fromkeys(to_compute))
|
||||||
|
prompts = [self._build_rag_prompt(s, self._retrieve_topk(s)) for s in uniq]
|
||||||
|
with torch.no_grad():
|
||||||
|
feat = self._rag_encode_batch(prompts, device)
|
||||||
|
for s, v in zip(uniq, feat):
|
||||||
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]
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user