微调T5文本纠错模型处理评论时部分样本无输出的解决方法
问题描述
我使用文本到文本模型T5对评论数据集进行拼写检查,分两轮微调,分别使用2万条和4万条评论,训练loss为0.0003,验证loss为0.000052。用20条评论样本手动测试时模型表现尚可,但处理超过1.4万条评论的数据集时,有1.1k条评论无输出。
数据集定义代码
class ReviewDataset(Dataset): def __init__(self, texts): self.inputs = ["fix: " + text for text in texts] def __len__(self): return len(self.inputs) def __getitem__(self, idx): return self.inputs[idx] def collate_fn(batch): encodings = tokenizer( batch, padding=True, truncation=True, max_length=128, return_tensors="pt" ) return encodings
批量处理代码
df = pd.read_csv("reviews.csv",encoding="latin-1") dataset = ReviewDataset(df["review_text"].tolist()) dataloader = DataLoader(dataset, batch_size=256, collate_fn=collate_fn) all_predictions = [] with torch.no_grad(): for batch in dataloader: input_ids = batch["input_ids"].to(device) attention_mask = batch["attention_mask"].to(device) outputs = model.generate(input_ids=input_ids, attention_mask=attention_mask, max_length=64) decoded = tokenizer.batch_decode(outputs, skip_special_tokens=True) all_predictions.extend(decoded) df["corrected_review"] = all_predictions df.to_csv("corrected_reviews_batched.csv", index=False)
解决方案
排查空输入与异常文本:先定位那1.1k条无输出的原始评论,检查是否存在空字符串、全特殊字符、极端超长文本等情况。可以在数据集初始化阶段过滤无效文本:
self.inputs = ["fix: " + text for text in texts if text.strip()]同时建议在原始DataFrame中保留索引,方便后续对应问题样本。
调整模型生成参数:T5默认的生成策略可能因提前触发终止符导致空输出,可添加以下参数优化:
min_length=1:强制生成至少1个有效tokenearly_stopping=False:避免过早停止生成num_beams=2:使用beam search提升生成稳定性
修改后的生成代码:
outputs = model.generate( input_ids=input_ids, attention_mask=attention_mask, max_length=64, min_length=1, early_stopping=False, num_beams=2 )检查tokenizer解码逻辑:
skip_special_tokens=True可能会过滤掉所有内容(比如模型仅输出了终止符)。可以临时关闭该参数查看原始输出,确认模型实际生成内容:decoded = tokenizer.batch_decode(outputs, skip_special_tokens=False)如果发现大量仅含
<pad>或<eos>的输出,说明这类输入不在模型的泛化范围内,需补充对应场景的微调数据。保障批量处理的稳定性:GPU内存不足可能导致部分batch静默失败,可尝试减小batch_size(如从256降至128),同时在循环中加入异常捕获,避免丢失数据对应关系:
with torch.no_grad(): for batch_idx, batch in enumerate(dataloader): try: input_ids = batch["input_ids"].to(device) attention_mask = batch["attention_mask"].to(device) outputs = model.generate(input_ids=input_ids, attention_mask=attention_mask, max_length=64, min_length=1) decoded = tokenizer.batch_decode(outputs, skip_special_tokens=True) all_predictions.extend(decoded) except Exception as e: print(f"Batch {batch_idx} processing failed: {str(e)}") # 为失败批次填充标记值,保证与原始数据行数一致 all_predictions.extend(["[ERROR]"] * len(batch["input_ids"]))对齐数据分布:若问题样本与微调数据分布差异过大(如含罕见语种、乱码、小众术语),需统计这类样本特征,补充到微调数据集重新训练,或在推理前对输入做预处理(如替换乱码、截断超长文本)。
内容的提问来源于stack exchange,提问作者Anurag Pandey
相关产品推荐
相关产品推荐

