From 724814851594dbef7389807a175a2f3d935150e6 Mon Sep 17 00:00:00 2001 From: Michelle0474 <2170308303@qq.com> Date: Wed, 1 Jul 2026 15:52:20 +0800 Subject: [PATCH] =?UTF-8?q?feat(llm):=20Qwen=20QLoRA=204-bit=20=E7=9C=9F?= =?UTF-8?q?=E5=BE=AE=E8=B0=83=E8=B7=AF=E5=BE=84=20+=20RAG=20=E7=BC=96?= =?UTF-8?q?=E7=A0=81=E5=86=BB=E7=BB=93/=E5=BE=AE=E8=B0=83=E5=88=86?= =?UTF-8?q?=E6=B5=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 显著提升 --- lnp_ml/modeling/layers/llm_prompt.py | 76 ++++++++++++++++++---------- 1 file changed, 48 insertions(+), 28 deletions(-) diff --git a/lnp_ml/modeling/layers/llm_prompt.py b/lnp_ml/modeling/layers/llm_prompt.py index 824dfef..7500873 100644 --- a/lnp_ml/modeling/layers/llm_prompt.py +++ b/lnp_ml/modeling/layers/llm_prompt.py @@ -54,8 +54,19 @@ class LLMPromptEncoder(nn.Module): self.encoder = T5EncoderModel.from_pretrained(model_name_or_path) self.hidden_size = self.encoder.config.d_model elif _is_qwen: - self.encoder = AutoModel.from_pretrained( - model_name_or_path, trust_remote_code=True, torch_dtype=torch.float16) + # 微调(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( + model_name_or_path, trust_remote_code=True, torch_dtype=torch.float16) self.hidden_size = self.encoder.config.hidden_size else: 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: 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(): p.requires_grad = False # 不同架构注意力层命名不同: @@ -153,34 +168,39 @@ class LLMPromptEncoder(nn.Module): "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): - """RAG 编码:每个分子→检索→构造prompt→Qwen编码→取最后有效token。带缓存。""" - feats = [] - 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] - 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): + """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() - 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: m = mask.unsqueeze(-1).float() # [B, L, 1]