mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-07-22 05:49:55 +08:00
Compare commits
No commits in common. "987ecf6386e505ee747642ce972f2c01fa25b4dc" and "0fc72a9474ab5fcb137599e2aa058e4b34373aae" have entirely different histories.
987ecf6386
...
0fc72a9474
@ -1,156 +1,152 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from typing import Dict, List, Literal, Optional, Tuple, Union
|
||||
|
||||
|
||||
PoolingStrategy = Literal["concat", "avg", "max", "attention"]
|
||||
|
||||
|
||||
class FusionLayer(nn.Module):
|
||||
"""
|
||||
将多个 token 融合成单个向量。
|
||||
|
||||
输入: Dict[str, Tensor] 或 [B, n_tokens, d_model]
|
||||
输出: [B, fusion_dim]
|
||||
|
||||
策略:
|
||||
- concat: [B, n_tokens, d_model] -> [B, n_tokens * d_model]
|
||||
- avg: [B, n_tokens, d_model] -> [B, d_model]
|
||||
- max: [B, n_tokens, d_model] -> [B, d_model]
|
||||
- attention: [B, n_tokens, d_model] -> [B, d_model] (learnable attention pooling)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_tokens: int,
|
||||
strategy: PoolingStrategy = "attention",
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
d_model: 每个 token 的维度
|
||||
n_tokens: token 数量(如 8)
|
||||
strategy: 融合策略
|
||||
"""
|
||||
super().__init__()
|
||||
self.d_model = d_model
|
||||
self.n_tokens = n_tokens
|
||||
self.strategy = strategy
|
||||
|
||||
if strategy == "concat":
|
||||
self.fusion_dim = n_tokens * d_model
|
||||
else:
|
||||
self.fusion_dim = d_model
|
||||
|
||||
# Attention pooling: learnable query
|
||||
if strategy == "attention":
|
||||
self.attn_query = nn.Parameter(torch.randn(1, 1, d_model))
|
||||
self.attn_proj = nn.Linear(d_model, d_model)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Union[Dict[str, torch.Tensor], torch.Tensor],
|
||||
return_attn_weights: bool = False,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""
|
||||
Args:
|
||||
x: Dict[str, Tensor] 每个 [B, d_model],或已 stack 的 [B, n_tokens, d_model]
|
||||
return_attn_weights: 若为 True 且策略为 attention,额外返回 attn_weights [B, n_tokens]
|
||||
|
||||
Returns:
|
||||
return_attn_weights=False: [B, fusion_dim]
|
||||
return_attn_weights=True: ([B, fusion_dim], [B, n_tokens])
|
||||
"""
|
||||
if isinstance(x, dict):
|
||||
x = torch.stack(list(x.values()), dim=1)
|
||||
|
||||
if self.strategy == "concat":
|
||||
out = x.flatten(start_dim=1)
|
||||
return (out, None) if return_attn_weights else out
|
||||
|
||||
elif self.strategy == "avg":
|
||||
out = x.mean(dim=1)
|
||||
return (out, None) if return_attn_weights else out
|
||||
|
||||
elif self.strategy == "max":
|
||||
out = x.max(dim=1).values
|
||||
return (out, None) if return_attn_weights else out
|
||||
|
||||
elif self.strategy == "attention":
|
||||
return self._attention_pooling(x, return_attn_weights)
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unknown strategy: {self.strategy}")
|
||||
|
||||
def _attention_pooling(
|
||||
self, x: torch.Tensor, return_attn_weights: bool = False,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""
|
||||
Attention pooling: 用可学习 query 对 tokens 做加权求和
|
||||
|
||||
Args:
|
||||
x: [B, n_tokens, d_model]
|
||||
return_attn_weights: 是否返回权重
|
||||
|
||||
Returns:
|
||||
return_attn_weights=False: [B, d_model]
|
||||
return_attn_weights=True: ([B, d_model], [B, n_tokens])
|
||||
"""
|
||||
B = x.size(0)
|
||||
query = self.attn_query.expand(B, -1, -1)
|
||||
|
||||
keys = self.attn_proj(x)
|
||||
scores = torch.bmm(query, keys.transpose(1, 2)) / (self.d_model ** 0.5)
|
||||
attn_weights = F.softmax(scores, dim=-1) # [B, 1, n_tokens]
|
||||
|
||||
out = torch.bmm(attn_weights, x).squeeze(1) # [B, d_model]
|
||||
|
||||
if return_attn_weights:
|
||||
return out, attn_weights.squeeze(1) # [B, n_tokens]
|
||||
return out
|
||||
|
||||
|
||||
class ResidualConcatFusion(nn.Module):
|
||||
"""对真实 token 做 attention pooling,再用零初始化门把 MoE/LLM 旁路以残差方式加入。
|
||||
|
||||
g_moe / g_llm 初始为 0 → +moe/+llm 起点严格等于 baseline;
|
||||
旁路只有确实有用时才会被训练打开,从机制上保证“加了不会更差”。
|
||||
"""
|
||||
|
||||
def __init__(self, d_model: int, strategy: PoolingStrategy = "attention") -> None:
|
||||
super().__init__()
|
||||
if strategy == "concat":
|
||||
raise ValueError("ResidualConcatFusion 不支持 concat(token 数随开关变化)")
|
||||
self.d_model = d_model
|
||||
self.pool = FusionLayer(d_model=d_model, n_tokens=1, strategy=strategy)
|
||||
self.fusion_dim = self.pool.fusion_dim
|
||||
# 零初始化门控(可学习标量),旁路初始不参与
|
||||
self.g_moe = nn.Parameter(torch.zeros(()))
|
||||
self.g_llm = nn.Parameter(torch.zeros(()))
|
||||
self.g_retr = nn.Parameter(torch.zeros(())) # 检索旁路零初始化门控
|
||||
|
||||
def forward(
|
||||
self,
|
||||
chem: torch.Tensor,
|
||||
tab: torch.Tensor,
|
||||
f_moe: Optional[torch.Tensor] = None,
|
||||
f_llm: Optional[torch.Tensor] = None,
|
||||
f_retr: Optional[torch.Tensor] = None,
|
||||
return_attn_weights: bool = False,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
# 只对真实 token(chem + tab)做注意力池化,旁路不参与 softmax 竞争
|
||||
seq = torch.cat([chem, tab], dim=1) # [B, n_chem + n_cond, d_model]
|
||||
pooled = self.pool(seq, return_attn_weights=return_attn_weights)
|
||||
if return_attn_weights:
|
||||
pooled, attn = pooled
|
||||
|
||||
out = pooled
|
||||
if f_moe is not None:
|
||||
out = out + self.g_moe * f_moe # 残差 + 零初始化门
|
||||
if f_llm is not None:
|
||||
out = out + self.g_llm * f_llm
|
||||
if f_retr is not None:
|
||||
out = out + self.g_retr * f_retr # 检索旁路,零初始化门保证起点=不开
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from typing import Dict, List, Literal, Optional, Tuple, Union
|
||||
|
||||
|
||||
PoolingStrategy = Literal["concat", "avg", "max", "attention"]
|
||||
|
||||
|
||||
class FusionLayer(nn.Module):
|
||||
"""
|
||||
将多个 token 融合成单个向量。
|
||||
|
||||
输入: Dict[str, Tensor] 或 [B, n_tokens, d_model]
|
||||
输出: [B, fusion_dim]
|
||||
|
||||
策略:
|
||||
- concat: [B, n_tokens, d_model] -> [B, n_tokens * d_model]
|
||||
- avg: [B, n_tokens, d_model] -> [B, d_model]
|
||||
- max: [B, n_tokens, d_model] -> [B, d_model]
|
||||
- attention: [B, n_tokens, d_model] -> [B, d_model] (learnable attention pooling)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_tokens: int,
|
||||
strategy: PoolingStrategy = "attention",
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
d_model: 每个 token 的维度
|
||||
n_tokens: token 数量(如 8)
|
||||
strategy: 融合策略
|
||||
"""
|
||||
super().__init__()
|
||||
self.d_model = d_model
|
||||
self.n_tokens = n_tokens
|
||||
self.strategy = strategy
|
||||
|
||||
if strategy == "concat":
|
||||
self.fusion_dim = n_tokens * d_model
|
||||
else:
|
||||
self.fusion_dim = d_model
|
||||
|
||||
# Attention pooling: learnable query
|
||||
if strategy == "attention":
|
||||
self.attn_query = nn.Parameter(torch.randn(1, 1, d_model))
|
||||
self.attn_proj = nn.Linear(d_model, d_model)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Union[Dict[str, torch.Tensor], torch.Tensor],
|
||||
return_attn_weights: bool = False,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""
|
||||
Args:
|
||||
x: Dict[str, Tensor] 每个 [B, d_model],或已 stack 的 [B, n_tokens, d_model]
|
||||
return_attn_weights: 若为 True 且策略为 attention,额外返回 attn_weights [B, n_tokens]
|
||||
|
||||
Returns:
|
||||
return_attn_weights=False: [B, fusion_dim]
|
||||
return_attn_weights=True: ([B, fusion_dim], [B, n_tokens])
|
||||
"""
|
||||
if isinstance(x, dict):
|
||||
x = torch.stack(list(x.values()), dim=1)
|
||||
|
||||
if self.strategy == "concat":
|
||||
out = x.flatten(start_dim=1)
|
||||
return (out, None) if return_attn_weights else out
|
||||
|
||||
elif self.strategy == "avg":
|
||||
out = x.mean(dim=1)
|
||||
return (out, None) if return_attn_weights else out
|
||||
|
||||
elif self.strategy == "max":
|
||||
out = x.max(dim=1).values
|
||||
return (out, None) if return_attn_weights else out
|
||||
|
||||
elif self.strategy == "attention":
|
||||
return self._attention_pooling(x, return_attn_weights)
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unknown strategy: {self.strategy}")
|
||||
|
||||
def _attention_pooling(
|
||||
self, x: torch.Tensor, return_attn_weights: bool = False,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""
|
||||
Attention pooling: 用可学习 query 对 tokens 做加权求和
|
||||
|
||||
Args:
|
||||
x: [B, n_tokens, d_model]
|
||||
return_attn_weights: 是否返回权重
|
||||
|
||||
Returns:
|
||||
return_attn_weights=False: [B, d_model]
|
||||
return_attn_weights=True: ([B, d_model], [B, n_tokens])
|
||||
"""
|
||||
B = x.size(0)
|
||||
query = self.attn_query.expand(B, -1, -1)
|
||||
|
||||
keys = self.attn_proj(x)
|
||||
scores = torch.bmm(query, keys.transpose(1, 2)) / (self.d_model ** 0.5)
|
||||
attn_weights = F.softmax(scores, dim=-1) # [B, 1, n_tokens]
|
||||
|
||||
out = torch.bmm(attn_weights, x).squeeze(1) # [B, d_model]
|
||||
|
||||
if return_attn_weights:
|
||||
return out, attn_weights.squeeze(1) # [B, n_tokens]
|
||||
return out
|
||||
|
||||
|
||||
class ResidualConcatFusion(nn.Module):
|
||||
"""对真实 token 做 attention pooling,再用零初始化门把 MoE/LLM 旁路以残差方式加入。
|
||||
|
||||
g_moe / g_llm 初始为 0 → +moe/+llm 起点严格等于 baseline;
|
||||
旁路只有确实有用时才会被训练打开,从机制上保证“加了不会更差”。
|
||||
"""
|
||||
|
||||
def __init__(self, d_model: int, strategy: PoolingStrategy = "attention") -> None:
|
||||
super().__init__()
|
||||
if strategy == "concat":
|
||||
raise ValueError("ResidualConcatFusion 不支持 concat(token 数随开关变化)")
|
||||
self.d_model = d_model
|
||||
self.pool = FusionLayer(d_model=d_model, n_tokens=1, strategy=strategy)
|
||||
self.fusion_dim = self.pool.fusion_dim
|
||||
# 零初始化门控(可学习标量),旁路初始不参与
|
||||
self.g_moe = nn.Parameter(torch.zeros(()))
|
||||
self.g_llm = nn.Parameter(torch.zeros(()))
|
||||
|
||||
def forward(
|
||||
self,
|
||||
chem: torch.Tensor,
|
||||
tab: torch.Tensor,
|
||||
f_moe: Optional[torch.Tensor] = None,
|
||||
f_llm: Optional[torch.Tensor] = None,
|
||||
return_attn_weights: bool = False,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
# 只对真实 token(chem + tab)做注意力池化,旁路不参与 softmax 竞争
|
||||
seq = torch.cat([chem, tab], dim=1) # [B, n_chem + n_cond, d_model]
|
||||
pooled = self.pool(seq, return_attn_weights=return_attn_weights)
|
||||
if return_attn_weights:
|
||||
pooled, attn = pooled
|
||||
|
||||
out = pooled
|
||||
if f_moe is not None:
|
||||
out = out + self.g_moe * f_moe # 残差 + 零初始化门
|
||||
if f_llm is not None:
|
||||
out = out + self.g_llm * f_llm
|
||||
|
||||
return (out, attn) if return_attn_weights else out
|
||||
@ -54,19 +54,8 @@ class LLMPromptEncoder(nn.Module):
|
||||
self.encoder = T5EncoderModel.from_pretrained(model_name_or_path)
|
||||
self.hidden_size = self.encoder.config.d_model
|
||||
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(
|
||||
model_name_or_path, trust_remote_code=True, torch_dtype=torch.float16)
|
||||
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)
|
||||
@ -93,10 +82,6 @@ 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
|
||||
# 不同架构注意力层命名不同:
|
||||
@ -168,39 +153,34 @@ class LLMPromptEncoder(nn.Module):
|
||||
"into your internal representation."
|
||||
)
|
||||
|
||||
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)
|
||||
|
||||
@torch.no_grad()
|
||||
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):
|
||||
"""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):
|
||||
self._cache[f"RAG::{self._rag_pool_id}::{s}"] = v.float().cpu()
|
||||
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)
|
||||
return torch.stack([self._cache[k] for k in keys]).to(device)
|
||||
|
||||
def _mean_pool(self, last_hidden: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
||||
m = mask.unsqueeze(-1).float() # [B, L, 1]
|
||||
|
||||
@ -1,229 +1,229 @@
|
||||
"""
|
||||
MoE (Mixture-of-Experts) 模块:sample-level + 跨模态路由。
|
||||
|
||||
设计要点(与现有 8-token 架构对齐):
|
||||
- Router 输入:tab token 池化后的 [B, d](即配方/实验条件向量)。
|
||||
- Expert 输入:chem token flatten 后的 [B, T_chem * d](化学侧 token)。
|
||||
- 路由粒度:每个样本只算一次 router → gates [B, K]。
|
||||
- Top-k 稀疏激活 + 可选训练态 jitter 噪声。
|
||||
- 返回 load-balancing aux loss,由 trainer 端按权重加进总 loss。
|
||||
|
||||
输出:
|
||||
F_moe: [B, d_model] 与单个 token 同维度,便于追加到 fusion 序列中。
|
||||
extras: dict 诊断与监控信息(aux loss、gates 等)。
|
||||
"""
|
||||
|
||||
from typing import Dict, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class MoEAttentionPool(nn.Module):
|
||||
"""与 FusionLayer 同款的 attention pooling,参数独立。
|
||||
|
||||
将一组 token [B, T, d] 池化成单个查询向量 [B, d],用作 router 的输入。
|
||||
"""
|
||||
|
||||
def __init__(self, d_model: int) -> None:
|
||||
super().__init__()
|
||||
self.d_model = d_model
|
||||
self.query = nn.Parameter(torch.randn(1, 1, d_model))
|
||||
self.proj = nn.Linear(d_model, d_model)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
x: [B, T, d_model]
|
||||
|
||||
Returns:
|
||||
[B, d_model]
|
||||
"""
|
||||
B = x.size(0)
|
||||
q = self.query.expand(B, -1, -1) # [B, 1, d]
|
||||
k = self.proj(x) # [B, T, d]
|
||||
scores = torch.bmm(q, k.transpose(1, 2)) / (self.d_model ** 0.5)
|
||||
weights = F.softmax(scores, dim=-1) # [B, 1, T]
|
||||
return torch.bmm(weights, x).squeeze(1) # [B, d]
|
||||
|
||||
|
||||
class MoERouter(nn.Module):
|
||||
"""单层 softmax router + Top-k 稀疏激活。
|
||||
|
||||
Args:
|
||||
d_model: router 输入维度。
|
||||
n_experts: 专家数量 K。
|
||||
top_k: 每个样本激活的专家数。
|
||||
jitter_noise: 训练态加在 logits 上的均匀噪声幅度,用于探索;评估态自动关闭。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_experts: int,
|
||||
top_k: int = 2,
|
||||
jitter_noise: float = 0.0,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
if not 1 <= top_k <= n_experts:
|
||||
raise ValueError(f"top_k 必须在 [1, n_experts={n_experts}],收到 {top_k}")
|
||||
self.n_experts = n_experts
|
||||
self.top_k = top_k
|
||||
self.jitter_noise = jitter_noise
|
||||
self.linear = nn.Linear(d_model, n_experts)
|
||||
|
||||
def forward(
|
||||
self, q: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Args:
|
||||
q: [B, d_model] 路由查询向量。
|
||||
|
||||
Returns:
|
||||
gates: [B, n_experts] top-k 后再归一化的稀疏概率。
|
||||
probs_full: [B, n_experts] 未掩码的 softmax 概率(用于 aux loss)。
|
||||
expert_mask: [B, n_experts] 0/1 掩码,标记被激活的专家。
|
||||
"""
|
||||
logits = self.linear(q) # [B, K]
|
||||
if self.training and self.jitter_noise > 0.0:
|
||||
noise = (torch.rand_like(logits) - 0.5) * self.jitter_noise
|
||||
logits = logits + noise
|
||||
|
||||
probs_full = F.softmax(logits, dim=-1) # [B, K]
|
||||
|
||||
topk_vals, topk_idx = probs_full.topk(self.top_k, dim=-1) # [B, k]
|
||||
expert_mask = torch.zeros_like(probs_full)
|
||||
expert_mask.scatter_(1, topk_idx, 1.0)
|
||||
|
||||
gates = probs_full * expert_mask
|
||||
gates = gates / gates.sum(dim=-1, keepdim=True).clamp(min=1e-9)
|
||||
return gates, probs_full, expert_mask
|
||||
|
||||
|
||||
class MoEExpert(nn.Module):
|
||||
"""单个专家 MLP:in_dim → hidden_dim → out_dim。"""
|
||||
|
||||
def __init__(
|
||||
self, in_dim: int, hidden_dim: int, out_dim: int, dropout: float = 0.1,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.net = nn.Sequential(
|
||||
nn.Linear(in_dim, hidden_dim),
|
||||
nn.GELU(),
|
||||
nn.Dropout(dropout),
|
||||
nn.Linear(hidden_dim, out_dim),
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.net(x)
|
||||
|
||||
|
||||
class MoEBlock(nn.Module):
|
||||
"""
|
||||
Sample-level 跨模态 MoE。
|
||||
|
||||
流程:
|
||||
1. tab tokens [B, T_tab, d] -> attention pool -> q_tab [B, d]
|
||||
2. q_tab -> Router -> gates [B, K] (+ probs_full / expert_mask)
|
||||
3. chem tokens [B, T_chem, d] -> flatten -> [B, T_chem * d]
|
||||
4. K 个 Expert MLP 并行处理 -> stack [B, K, d]
|
||||
5. 加权求和 -> F_moe [B, d]
|
||||
6. 计算 load-balancing aux loss
|
||||
|
||||
Args:
|
||||
d_model: token 维度。
|
||||
n_chem_tokens: chem 侧 token 数(决定 expert 输入维度)。
|
||||
n_experts: 专家数量 K。
|
||||
top_k: 每个样本激活的专家数。
|
||||
expert_hidden_mult: expert 中间层维度 = expert_hidden_mult * d_model。
|
||||
dropout: expert 内部 dropout。
|
||||
jitter_noise: router 训练态噪声幅度,0.0 表示关闭。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_chem_tokens: int = 4,
|
||||
n_experts: int = 4,
|
||||
top_k: int = 2,
|
||||
expert_hidden_mult: int = 2,
|
||||
dropout: float = 0.1,
|
||||
jitter_noise: float = 0.0,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.d_model = d_model
|
||||
self.n_chem_tokens = n_chem_tokens
|
||||
self.n_experts = n_experts
|
||||
self.top_k = top_k
|
||||
|
||||
self.tab_pool = MoEAttentionPool(d_model)
|
||||
self.router = MoERouter(
|
||||
d_model, n_experts, top_k=top_k, jitter_noise=jitter_noise,
|
||||
)
|
||||
|
||||
expert_in = d_model * n_chem_tokens
|
||||
expert_hidden = d_model * expert_hidden_mult
|
||||
self.experts = nn.ModuleList([
|
||||
MoEExpert(expert_in, expert_hidden, d_model, dropout=dropout)
|
||||
for _ in range(n_experts)
|
||||
])
|
||||
|
||||
def forward(
|
||||
self,
|
||||
chem: torch.Tensor,
|
||||
tab: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
"""
|
||||
Args:
|
||||
chem: [B, T_chem, d_model] 化学侧 token(router 不看)。
|
||||
tab: [B, T_tab, d_model] 配方/实验侧 token(router 看)。
|
||||
|
||||
Returns:
|
||||
F_moe: [B, d_model]
|
||||
extras: {
|
||||
"lb_loss": 标量 load-balancing aux loss(带梯度),
|
||||
"gates": [B, K] 稀疏归一后的门控(detached),
|
||||
"probs": [B, K] 原始 softmax 概率(detached),
|
||||
}
|
||||
"""
|
||||
if chem.size(1) != self.n_chem_tokens:
|
||||
raise ValueError(
|
||||
f"chem token 数不匹配:期望 {self.n_chem_tokens}, "
|
||||
f"实际 {chem.size(1)}"
|
||||
)
|
||||
|
||||
q_tab = self.tab_pool(tab) # [B, d]
|
||||
gates, probs_full, expert_mask = self.router(q_tab) # 各 [B, K]
|
||||
|
||||
flat = chem.flatten(start_dim=1) # [B, T_chem * d]
|
||||
expert_outs = torch.stack(
|
||||
[expert(flat) for expert in self.experts], dim=1,
|
||||
) # [B, K, d]
|
||||
|
||||
F_moe = (gates.unsqueeze(-1) * expert_outs).sum(dim=1) # [B, d]
|
||||
|
||||
lb_loss = self._load_balancing_loss(probs_full, expert_mask)
|
||||
|
||||
return F_moe, {
|
||||
"lb_loss": lb_loss,
|
||||
"gates": gates.detach(),
|
||||
"probs": probs_full.detach(),
|
||||
}
|
||||
|
||||
def _load_balancing_loss(
|
||||
self,
|
||||
probs_full: torch.Tensor,
|
||||
expert_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Switch Transformer 风格的 load-balancing loss。
|
||||
|
||||
f_i = 该 batch 中 expert_i 被激活的样本占比(top-k mask 的均值)
|
||||
p_i = 该 batch 中 expert_i 的 softmax 概率均值
|
||||
loss = K * Σ_i f_i * p_i
|
||||
|
||||
理想情况下 f_i 与 p_i 都接近 1/K,loss ≈ 1。
|
||||
"""
|
||||
f = expert_mask.mean(dim=0) # [K]
|
||||
p = probs_full.mean(dim=0) # [K]
|
||||
"""
|
||||
MoE (Mixture-of-Experts) 模块:sample-level + 跨模态路由。
|
||||
|
||||
设计要点(与现有 8-token 架构对齐):
|
||||
- Router 输入:tab token 池化后的 [B, d](即配方/实验条件向量)。
|
||||
- Expert 输入:chem token flatten 后的 [B, T_chem * d](化学侧 token)。
|
||||
- 路由粒度:每个样本只算一次 router → gates [B, K]。
|
||||
- Top-k 稀疏激活 + 可选训练态 jitter 噪声。
|
||||
- 返回 load-balancing aux loss,由 trainer 端按权重加进总 loss。
|
||||
|
||||
输出:
|
||||
F_moe: [B, d_model] 与单个 token 同维度,便于追加到 fusion 序列中。
|
||||
extras: dict 诊断与监控信息(aux loss、gates 等)。
|
||||
"""
|
||||
|
||||
from typing import Dict, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class MoEAttentionPool(nn.Module):
|
||||
"""与 FusionLayer 同款的 attention pooling,参数独立。
|
||||
|
||||
将一组 token [B, T, d] 池化成单个查询向量 [B, d],用作 router 的输入。
|
||||
"""
|
||||
|
||||
def __init__(self, d_model: int) -> None:
|
||||
super().__init__()
|
||||
self.d_model = d_model
|
||||
self.query = nn.Parameter(torch.randn(1, 1, d_model))
|
||||
self.proj = nn.Linear(d_model, d_model)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
x: [B, T, d_model]
|
||||
|
||||
Returns:
|
||||
[B, d_model]
|
||||
"""
|
||||
B = x.size(0)
|
||||
q = self.query.expand(B, -1, -1) # [B, 1, d]
|
||||
k = self.proj(x) # [B, T, d]
|
||||
scores = torch.bmm(q, k.transpose(1, 2)) / (self.d_model ** 0.5)
|
||||
weights = F.softmax(scores, dim=-1) # [B, 1, T]
|
||||
return torch.bmm(weights, x).squeeze(1) # [B, d]
|
||||
|
||||
|
||||
class MoERouter(nn.Module):
|
||||
"""单层 softmax router + Top-k 稀疏激活。
|
||||
|
||||
Args:
|
||||
d_model: router 输入维度。
|
||||
n_experts: 专家数量 K。
|
||||
top_k: 每个样本激活的专家数。
|
||||
jitter_noise: 训练态加在 logits 上的均匀噪声幅度,用于探索;评估态自动关闭。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_experts: int,
|
||||
top_k: int = 2,
|
||||
jitter_noise: float = 0.0,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
if not 1 <= top_k <= n_experts:
|
||||
raise ValueError(f"top_k 必须在 [1, n_experts={n_experts}],收到 {top_k}")
|
||||
self.n_experts = n_experts
|
||||
self.top_k = top_k
|
||||
self.jitter_noise = jitter_noise
|
||||
self.linear = nn.Linear(d_model, n_experts)
|
||||
|
||||
def forward(
|
||||
self, q: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Args:
|
||||
q: [B, d_model] 路由查询向量。
|
||||
|
||||
Returns:
|
||||
gates: [B, n_experts] top-k 后再归一化的稀疏概率。
|
||||
probs_full: [B, n_experts] 未掩码的 softmax 概率(用于 aux loss)。
|
||||
expert_mask: [B, n_experts] 0/1 掩码,标记被激活的专家。
|
||||
"""
|
||||
logits = self.linear(q) # [B, K]
|
||||
if self.training and self.jitter_noise > 0.0:
|
||||
noise = (torch.rand_like(logits) - 0.5) * self.jitter_noise
|
||||
logits = logits + noise
|
||||
|
||||
probs_full = F.softmax(logits, dim=-1) # [B, K]
|
||||
|
||||
topk_vals, topk_idx = probs_full.topk(self.top_k, dim=-1) # [B, k]
|
||||
expert_mask = torch.zeros_like(probs_full)
|
||||
expert_mask.scatter_(1, topk_idx, 1.0)
|
||||
|
||||
gates = probs_full * expert_mask
|
||||
gates = gates / gates.sum(dim=-1, keepdim=True).clamp(min=1e-9)
|
||||
return gates, probs_full, expert_mask
|
||||
|
||||
|
||||
class MoEExpert(nn.Module):
|
||||
"""单个专家 MLP:in_dim → hidden_dim → out_dim。"""
|
||||
|
||||
def __init__(
|
||||
self, in_dim: int, hidden_dim: int, out_dim: int, dropout: float = 0.1,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.net = nn.Sequential(
|
||||
nn.Linear(in_dim, hidden_dim),
|
||||
nn.GELU(),
|
||||
nn.Dropout(dropout),
|
||||
nn.Linear(hidden_dim, out_dim),
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.net(x)
|
||||
|
||||
|
||||
class MoEBlock(nn.Module):
|
||||
"""
|
||||
Sample-level 跨模态 MoE。
|
||||
|
||||
流程:
|
||||
1. tab tokens [B, T_tab, d] -> attention pool -> q_tab [B, d]
|
||||
2. q_tab -> Router -> gates [B, K] (+ probs_full / expert_mask)
|
||||
3. chem tokens [B, T_chem, d] -> flatten -> [B, T_chem * d]
|
||||
4. K 个 Expert MLP 并行处理 -> stack [B, K, d]
|
||||
5. 加权求和 -> F_moe [B, d]
|
||||
6. 计算 load-balancing aux loss
|
||||
|
||||
Args:
|
||||
d_model: token 维度。
|
||||
n_chem_tokens: chem 侧 token 数(决定 expert 输入维度)。
|
||||
n_experts: 专家数量 K。
|
||||
top_k: 每个样本激活的专家数。
|
||||
expert_hidden_mult: expert 中间层维度 = expert_hidden_mult * d_model。
|
||||
dropout: expert 内部 dropout。
|
||||
jitter_noise: router 训练态噪声幅度,0.0 表示关闭。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_chem_tokens: int = 4,
|
||||
n_experts: int = 4,
|
||||
top_k: int = 2,
|
||||
expert_hidden_mult: int = 2,
|
||||
dropout: float = 0.1,
|
||||
jitter_noise: float = 0.0,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.d_model = d_model
|
||||
self.n_chem_tokens = n_chem_tokens
|
||||
self.n_experts = n_experts
|
||||
self.top_k = top_k
|
||||
|
||||
self.tab_pool = MoEAttentionPool(d_model)
|
||||
self.router = MoERouter(
|
||||
d_model, n_experts, top_k=top_k, jitter_noise=jitter_noise,
|
||||
)
|
||||
|
||||
expert_in = d_model * n_chem_tokens
|
||||
expert_hidden = d_model * expert_hidden_mult
|
||||
self.experts = nn.ModuleList([
|
||||
MoEExpert(expert_in, expert_hidden, d_model, dropout=dropout)
|
||||
for _ in range(n_experts)
|
||||
])
|
||||
|
||||
def forward(
|
||||
self,
|
||||
chem: torch.Tensor,
|
||||
tab: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
"""
|
||||
Args:
|
||||
chem: [B, T_chem, d_model] 化学侧 token(router 不看)。
|
||||
tab: [B, T_tab, d_model] 配方/实验侧 token(router 看)。
|
||||
|
||||
Returns:
|
||||
F_moe: [B, d_model]
|
||||
extras: {
|
||||
"lb_loss": 标量 load-balancing aux loss(带梯度),
|
||||
"gates": [B, K] 稀疏归一后的门控(detached),
|
||||
"probs": [B, K] 原始 softmax 概率(detached),
|
||||
}
|
||||
"""
|
||||
if chem.size(1) != self.n_chem_tokens:
|
||||
raise ValueError(
|
||||
f"chem token 数不匹配:期望 {self.n_chem_tokens}, "
|
||||
f"实际 {chem.size(1)}"
|
||||
)
|
||||
|
||||
q_tab = self.tab_pool(tab) # [B, d]
|
||||
gates, probs_full, expert_mask = self.router(q_tab) # 各 [B, K]
|
||||
|
||||
flat = chem.flatten(start_dim=1) # [B, T_chem * d]
|
||||
expert_outs = torch.stack(
|
||||
[expert(flat) for expert in self.experts], dim=1,
|
||||
) # [B, K, d]
|
||||
|
||||
F_moe = (gates.unsqueeze(-1) * expert_outs).sum(dim=1) # [B, d]
|
||||
|
||||
lb_loss = self._load_balancing_loss(probs_full, expert_mask)
|
||||
|
||||
return F_moe, {
|
||||
"lb_loss": lb_loss,
|
||||
"gates": gates.detach(),
|
||||
"probs": probs_full.detach(),
|
||||
}
|
||||
|
||||
def _load_balancing_loss(
|
||||
self,
|
||||
probs_full: torch.Tensor,
|
||||
expert_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Switch Transformer 风格的 load-balancing loss。
|
||||
|
||||
f_i = 该 batch 中 expert_i 被激活的样本占比(top-k mask 的均值)
|
||||
p_i = 该 batch 中 expert_i 的 softmax 概率均值
|
||||
loss = K * Σ_i f_i * p_i
|
||||
|
||||
理想情况下 f_i 与 p_i 都接近 1/K,loss ≈ 1。
|
||||
"""
|
||||
f = expert_mask.mean(dim=0) # [K]
|
||||
p = probs_full.mean(dim=0) # [K]
|
||||
return self.n_experts * (f * p).sum()
|
||||
@ -632,21 +632,6 @@ def _run_single_outer_fold(
|
||||
fold_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
full_dataset = LNPDataset(df)
|
||||
# === 断点续跑:已完成的 fold 直接跳过(读回磁盘结果)===
|
||||
_tm = fold_dir / "test_metrics.json"
|
||||
_bp = fold_dir / "best_params.json"
|
||||
_em = fold_dir / "epoch_mean.json"
|
||||
if precomputed_best_params is None and _tm.exists() and _bp.exists() and _em.exists():
|
||||
logger.success(f"[SKIP] Outer fold {outer_fold} already done, loading cached results.")
|
||||
with open(_tm) as _f: _tmd = json.load(_f)
|
||||
with open(_bp) as _f: _bpd = json.load(_f)
|
||||
with open(_em) as _f: _emd = json.load(_f)
|
||||
return {
|
||||
"fold": outer_fold,
|
||||
"best_params": _bpd,
|
||||
"epoch_mean": int(_emd.get("epoch_mean", _emd) if isinstance(_emd, dict) else _emd),
|
||||
"test_metrics": _tmd,
|
||||
}
|
||||
|
||||
logger.info(f"\n{'='*60}")
|
||||
logger.info(f"OUTER FOLD {outer_fold}")
|
||||
|
||||
@ -1,193 +0,0 @@
|
||||
"""
|
||||
检索增强特征模块 (第一层优化)。
|
||||
|
||||
核心思想 (借鉴 MolRAG 的化学硬先验「结构相似 → 性质相似」):
|
||||
对每个分子,用 Morgan 指纹在「训练集检索池」中找最相似的 k 个邻居,
|
||||
把它们的已知标签 (已 z-score 的 quantified_delivery 等) 聚合成一个特征向量,
|
||||
作为额外旁路注入融合层 (配零初始化门控,保证「加了不会更差」)。
|
||||
|
||||
防数据泄漏铁律 (本模块通过接口与逻辑强制保证):
|
||||
1. 检索池 (pool_smiles / pool_labels) 只能由调用方传入「训练集」分子。
|
||||
本模块自身不接触任何全局数据,只用传进来的池。
|
||||
2. exclude_self=True 时,查询分子若在池中 (训练分子查询自己),
|
||||
排除指纹完全相同的「自己」(相似度=1.0 的那一个),否则等于直接看答案。
|
||||
3. 测试分子查询时,池里全是训练集,测试分子不在池中,
|
||||
天然不会检索到自己或其他测试分子。
|
||||
|
||||
设计原则:
|
||||
- 纯特征工程,不引入可训练参数 (投影层在 models.py 侧)。
|
||||
- 池构建时一次性预计算所有指纹,查询时只算查询分子指纹 + 相似度,高效。
|
||||
- 标签已在 dataset.py 标准化 (z-score),聚合时直接 (加权) 平均即可。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List, Optional, Sequence
|
||||
|
||||
import numpy as np
|
||||
|
||||
from rdkit import Chem
|
||||
from rdkit.Chem import AllChem, DataStructs
|
||||
from rdkit import RDLogger
|
||||
|
||||
# 关闭 RDKit 的噪音日志 (无效 SMILES 警告等)
|
||||
RDLogger.DisableLog("rdApp.*")
|
||||
|
||||
|
||||
def _smiles_to_fp(smiles: str, radius: int = 2, n_bits: int = 2048):
|
||||
"""单个 SMILES -> Morgan 指纹 (ExplicitBitVect)。无效 SMILES 返回 None。"""
|
||||
mol = Chem.MolFromSmiles(smiles)
|
||||
if mol is None:
|
||||
return None
|
||||
return AllChem.GetMorganFingerprintAsBitVect(mol, radius=radius, nBits=n_bits)
|
||||
|
||||
|
||||
class MorganRetriever:
|
||||
"""基于 Morgan 指纹 + Tanimoto 相似度的分子检索器。
|
||||
|
||||
用法:
|
||||
# 只用训练集分子建池
|
||||
retr = MorganRetriever(
|
||||
pool_smiles=train_smiles, # List[str]
|
||||
pool_labels=train_labels, # np.ndarray [N_pool, label_dim]
|
||||
radius=2, n_bits=2048, k=5,
|
||||
)
|
||||
# 查询 (训练分子查询时 exclude_self=True 防泄漏)
|
||||
feat = retr.query(query_smiles, exclude_self=True) # np.ndarray [label_dim*?]
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pool_smiles: Sequence[str],
|
||||
pool_labels: np.ndarray,
|
||||
radius: int = 2,
|
||||
n_bits: int = 2048,
|
||||
k: int = 5,
|
||||
sim_weighted: bool = True,
|
||||
include_sim_stats: bool = True,
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
pool_smiles: 检索池分子的 SMILES (必须只含训练集分子)。
|
||||
pool_labels: 检索池分子的标签, shape [N_pool, label_dim], 已标准化。
|
||||
radius: Morgan 指纹半径 (MolRAG 用 2)。
|
||||
n_bits: 指纹位数。
|
||||
k: 检索近邻数。
|
||||
sim_weighted: 聚合邻居标签时是否按相似度加权 (否则等权平均)。
|
||||
include_sim_stats: 是否在输出特征里附带相似度统计 (top1/mean 相似度),
|
||||
让模型知道「邻居有多可靠」。
|
||||
"""
|
||||
assert len(pool_smiles) == len(pool_labels), "池 SMILES 与标签数量不一致"
|
||||
self.radius = radius
|
||||
self.n_bits = n_bits
|
||||
self.k = k
|
||||
self.sim_weighted = sim_weighted
|
||||
self.include_sim_stats = include_sim_stats
|
||||
|
||||
_smiles_arr = list(pool_smiles)
|
||||
_labels_arr = np.asarray(pool_labels, dtype=np.float32)
|
||||
# 过滤掉标签含 NaN 的分子:它们没有有效标签,不能作为性质参考邻居
|
||||
_valid_mask = ~np.isnan(_labels_arr).any(axis=1)
|
||||
self.pool_smiles: List[str] = [_smiles_arr[i] for i in range(len(_smiles_arr)) if _valid_mask[i]]
|
||||
self.pool_labels = _labels_arr[_valid_mask]
|
||||
self.label_dim = self.pool_labels.shape[1]
|
||||
_n_dropped = int((~_valid_mask).sum())
|
||||
if _n_dropped > 0:
|
||||
import sys
|
||||
print(f'[MorganRetriever] 过滤 {_n_dropped} 个 NaN 标签分子,有效检索池={len(self.pool_smiles)}', file=sys.stderr)
|
||||
|
||||
# 池内 SMILES 集合,用于 exclude_self="auto" 时判断查询分子是否在池中
|
||||
self.pool_smiles_set = set(self.pool_smiles)
|
||||
|
||||
# 预计算池内所有指纹 (无效分子记 None,检索时跳过)
|
||||
self.pool_fps = [_smiles_to_fp(s, radius, n_bits) for s in self.pool_smiles]
|
||||
|
||||
# 标签均值,用于无有效邻居时的回退 (用训练集均值,不泄漏)
|
||||
self.label_mean = self.pool_labels.mean(axis=0)
|
||||
|
||||
@property
|
||||
def feature_dim(self) -> int:
|
||||
"""输出特征维度: 聚合标签 (label_dim) [+ 相似度统计 2 维]。"""
|
||||
return self.label_dim + (2 if self.include_sim_stats else 0)
|
||||
|
||||
def query(self, query_smiles: str, exclude_self="auto") -> np.ndarray:
|
||||
"""检索查询分子的 top-k 邻居并聚合其标签为特征向量。
|
||||
|
||||
Args:
|
||||
query_smiles: 查询分子 SMILES。
|
||||
exclude_self: 排除自身策略,防数据泄漏的关键。
|
||||
- "auto" (默认, 推荐): 查询分子若在检索池中(训练分子),自动排除自己;
|
||||
不在池中(测试分子)则不排除。无需调用方区分训练/测试。
|
||||
- True: 强制排除相似度=1.0 的分子。
|
||||
- False: 不排除(仅当确定查询分子不在池中时使用)。
|
||||
|
||||
Returns:
|
||||
特征向量 np.ndarray [feature_dim]。无有效邻居时回退到训练集均值。
|
||||
"""
|
||||
# auto 模式:查询分子在池中 → 训练分子 → 排除自己;否则不排除
|
||||
if exclude_self == "auto":
|
||||
do_exclude = query_smiles in self.pool_smiles_set
|
||||
else:
|
||||
do_exclude = bool(exclude_self)
|
||||
|
||||
q_fp = _smiles_to_fp(query_smiles, self.radius, self.n_bits)
|
||||
if q_fp is None:
|
||||
# 查询分子无效:回退到训练集均值 + 零相似度
|
||||
agg = self.label_mean.copy()
|
||||
if self.include_sim_stats:
|
||||
agg = np.concatenate([agg, np.array([0.0, 0.0], dtype=np.float32)])
|
||||
return agg.astype(np.float32)
|
||||
|
||||
# 与池内每个分子算 Tanimoto 相似度
|
||||
sims = np.full(len(self.pool_fps), -1.0, dtype=np.float32)
|
||||
for i, fp in enumerate(self.pool_fps):
|
||||
if fp is None:
|
||||
continue
|
||||
sims[i] = DataStructs.TanimotoSimilarity(q_fp, fp)
|
||||
|
||||
# 排除自己: 指纹完全相同 (相似度 >= 1.0 - eps) 的那一个
|
||||
if do_exclude:
|
||||
eps = 1e-6
|
||||
self_mask = sims >= (1.0 - eps)
|
||||
# 保守起见:把所有相似度=1.0 的都视作潜在自身并排除,
|
||||
# 因为结构完全相同的分子标签也应相同,留着等于泄漏答案。
|
||||
sims[self_mask] = -1.0
|
||||
|
||||
# 取 top-k (相似度降序),过滤掉无效 (-1) 的
|
||||
valid = np.where(sims >= 0.0)[0]
|
||||
if len(valid) == 0:
|
||||
# 没有有效邻居:回退到训练集均值
|
||||
agg = self.label_mean.copy()
|
||||
if self.include_sim_stats:
|
||||
agg = np.concatenate([agg, np.array([0.0, 0.0], dtype=np.float32)])
|
||||
return agg.astype(np.float32)
|
||||
|
||||
order = valid[np.argsort(-sims[valid])]
|
||||
topk_idx = order[: self.k]
|
||||
topk_sims = sims[topk_idx]
|
||||
topk_labels = self.pool_labels[topk_idx] # [k', label_dim]
|
||||
|
||||
# 聚合邻居标签
|
||||
if self.sim_weighted and topk_sims.sum() > 1e-8:
|
||||
w = topk_sims / topk_sims.sum()
|
||||
agg_label = (w[:, None] * topk_labels).sum(axis=0)
|
||||
else:
|
||||
agg_label = topk_labels.mean(axis=0)
|
||||
|
||||
if self.include_sim_stats:
|
||||
sim_stats = np.array(
|
||||
[float(topk_sims[0]), float(topk_sims.mean())], dtype=np.float32
|
||||
)
|
||||
agg = np.concatenate([agg_label.astype(np.float32), sim_stats])
|
||||
else:
|
||||
agg = agg_label.astype(np.float32)
|
||||
|
||||
return agg.astype(np.float32)
|
||||
|
||||
def query_batch(
|
||||
self, query_smiles_list, exclude_self="auto"
|
||||
) -> np.ndarray:
|
||||
"""批量检索。返回 [N_query, feature_dim]。"""
|
||||
return np.stack(
|
||||
[self.query(s, exclude_self=exclude_self) for s in query_smiles_list]
|
||||
)
|
||||
Loading…
x
Reference in New Issue
Block a user