File size: 5,744 Bytes
f73f9b3 | 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 163 | import torch
import torch.nn as nn
class LightweightTextEncoder(nn.Module):
"""轻量级字符/词元编码器(离线 fallback)。
当无法下载预训练模型时,使用基于字符嵌入的简单编码器,
保证整个流程在离线 CPU 环境也能运行。
"""
def __init__(self, hidden_dim=128, vocab_dim=256, max_len=128):
super().__init__()
self.hidden_dim = hidden_dim
self.max_len = max_len
self.embedding = nn.Embedding(vocab_dim, hidden_dim)
self.encoder_layer = nn.TransformerEncoderLayer(
d_model=hidden_dim, nhead=4, dim_feedforward=hidden_dim * 2, batch_first=True
)
self.encoder_transformer = nn.TransformerEncoder(self.encoder_layer, num_layers=2)
def tokenize_texts(self, texts):
"""将文本转换为字符 ID 序列。"""
batch_sequences = []
for text in texts:
byte_ids = [min(b, 255) for b in text.encode("utf-8", errors="ignore")[: self.max_len]]
pad_len = self.max_len - len(byte_ids)
byte_ids = byte_ids + [0] * pad_len
batch_sequences.append(byte_ids)
return torch.tensor(batch_sequences, dtype=torch.long)
def forward(self, texts):
"""编码文本。
Args:
texts: List[str]
Returns:
embeddings: (batch, hidden_dim)
"""
device = next(self.parameters()).device
tokens = self.tokenize_texts(texts).to(device)
token_mask = (tokens != 0).unsqueeze(-1).float()
embedded = self.embedding(tokens)
embedded = embedded * token_mask
features = self.encoder_transformer(embedded)
pooled = features.mean(dim=1)
norm = torch.norm(pooled, p=2, dim=-1, keepdim=True).clamp(min=1e-6)
pooled = pooled / norm
return pooled
class LLMEncoder(nn.Module):
"""LLM 编码器封装。
将开源小模型(如 distilbert-base-uncased)作为语义编码器,
支持冻结或 LoRA 微调。若无法联网下载预训练模型,
会自动回退到内置的轻量级编码器(离线可用)。
"""
def __init__(self, config):
super().__init__()
self.encoder_name = config.get("encoder_name", "distilbert-base-uncased")
self.freeze = config.get("encoder_freeze", True)
self.use_lora = config.get("use_lora", True)
self.hidden_dim = config.get("hidden_dim", 768)
self.backend = "none"
self.encoder = self._build_encoder(config)
def _build_encoder(self, config):
import os
HF_HUB_OFFLINE = config.get("hf_offline", False)
if HF_HUB_OFFLINE:
return self._build_lightweight(config)
try:
from transformers import AutoModel, AutoTokenizer
from peft import LoraConfig, get_peft_model, TaskType
self.tokenizer = AutoTokenizer.from_pretrained(self.encoder_name)
encoder = AutoModel.from_pretrained(self.encoder_name)
if self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
if self.freeze:
for param in encoder.parameters():
param.requires_grad = False
if self.use_lora and not self.freeze:
lora_config = LoraConfig(
task_type=TaskType.FEATURE_EXTRACTION,
r=config.get("lora_r", 8),
lora_alpha=config.get("lora_alpha", 16),
lora_dropout=config.get("lora_dropout", 0.1),
target_modules=["q_lin", "k_lin", "v_lin", "out_lin"],
)
encoder = get_peft_model(encoder, lora_config)
self.output_projection = nn.Linear(
encoder.config.hidden_size, self.hidden_dim
)
self.backend = "transformers"
return encoder
except Exception as e:
print(
f"[Encoder] Failed to load pretrained model '{self.encoder_name}': {e}\n"
f"[Encoder] Falling back to lightweight local encoder (offline mode)."
)
return self._build_lightweight(config)
def _build_lightweight(self, config):
"""构建轻量级离线编码器。"""
embed_dim = config.get("embed_dim", min(128, self.hidden_dim))
self.backend = "lightweight"
encoder = LightweightTextEncoder(
hidden_dim=embed_dim, vocab_dim=256, max_len=128
)
self.output_projection = nn.Linear(embed_dim, self.hidden_dim)
self.encode_text = encoder.forward
return encoder
def forward(self, texts):
"""编码输入文本为语义嵌入。
Args:
texts: List[str] 输入文本
Returns:
embeddings: (batch, hidden_dim) 语义嵌入
"""
if self.backend == "lightweight":
raw = self._pylight_encode(texts)
return self.output_projection(raw)
inputs = self.tokenizer(
texts,
padding=True,
truncation=True,
max_length=128,
return_tensors="pt",
)
inputs = {k: v.to(self.encoder.device) for k, v in inputs.items()}
with torch.set_grad_enabled(not self.freeze):
outputs = self.encoder(**inputs)
embeddings = outputs.last_hidden_state[:, 0, :]
embeddings = self.output_projection(embeddings)
return embeddings
def _pylight_encode(self, texts):
"""内部轻量编码(保持 batch_first 等)。"""
if isinstance(texts, str):
texts = [texts]
return self.encoder(texts)
|