当 WeLM 在生成较长文本时直接报出“CUDA out of memory”,先不要急着改模型结构。这块错误通常只在显存分配失败时抛出,但分配发生在加载阶段还是推理阶段,处理方式完全不同。推理阶段如果显存是随生成长度逐步涨上去的,那么分块生成就是最直接的规避路径;如果加载阶段就已溢出,则要先减小模型加载占用或调整并行方式。
生成长文本时的显存溢出,大概率由单次生成长度引起,少数情况是加载预分配过大。先用缓存分配器日志和 nvidia-smi 输出定位阶段,再把一次性生成改为分块或滑动窗口生成,最后用生成完整性做验证。适用场景是单卡或固定显存环境;不换卡不做量化。风险是分块过长会破坏上下文连续,因此需保留重叠区;建议从短序列逐步加大长度观察显存曲线。
区分加载期与推理期的OOM报错
先看堆栈信息:在 Python 里捕获异常并打印堆栈,或用 PYTORCH_NO_CUDA_LAZY_LOADING=1 启动,让 CUDA 初始化更可观测。PyTorch 的缓存分配器日志可以直接看分配行为,运行程序前设置环境变量 PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,max_split_size_mb:128,并打开 debug 日志,像这样:
import torch
import traceback
try:
outputs = model.generate(input_ids, max_length=long_length)
except torch.cuda.OutOfMemoryError:
traceback.print_exc()
print("torch.cuda.memory_summary():")
print(torch.cuda.memory_summary(abbreviated=False))如果堆栈最后停在 model.generate 或对应的 forward 函数内,说明是推理期动态分配;如果堆栈停在模型加载或权重转移到 CUDA 的代码,则是加载期。同时打开另一个终端运行 nvidia-smi -l 1,每秒刷新一次。观察显存是稳定在一个高位,还是随生成步数线性增加。加载期通常几秒内冲到峰值后不再变化,推理期则是逐步爬升直到溢出。
减小单次生成长度并观察显存曲线
确认是推理期溢出后,先做一个“长度-显存”曲线。不要直接用目标长度跑,而是从较小的长度开始逐步加大,记录每次的峰值显存。下面是一个简单的循环脚本,它会尝试每个长度,并在成功或失败时把结果追加到 csv 里:
import csv
import torch
lengths = [128, 256, 512, 768, 1024]
csv_path = 'welm_oom_probe.csv'
fields = ['max_length', 'success', 'allocated_mb', 'reserved_mb']
with open(csv_path, mode='w', newline='') as f:
writer = csv.DictWriter(f, fieldnames=fields)
writer.writeheader()
for max_len in lengths:
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
try:
output = model.generate(input_ids, max_length=max_len)
torch.cuda.synchronize()
allocated = torch.cuda.max_memory_allocated() / (1024**2)
reserved = torch.cuda.max_memory_reserved() / (1024**2)
writer.writerow({'max_length': max_len, 'success': 'yes',
'allocated_mb': f'{allocated:.1f}',
'reserved_mb': f'{reserved:.1f}'})
except torch.cuda.OutOfMemoryError:
writer.writerow({'max_length': max_len, 'success': 'no',
'allocated_mb': 'oom', 'reserved_mb': 'oom'})
f.flush()
print('probe done')这里 max_length 指的是整个生成序列的最终长度,不是新 token 数量。如果 512 成功、768 失败,那么边界就在这个区间内。记录下成功时的最大长度,作为后续分块生成的上限参考。脚本里用了 empty_cache(),是为避免上一次分配影响下一次测量。
实现上下文分块与拼接策略
当单次生成长度到不了目标长度时,不换卡的情况下可以分块。这里说的分块是把已生成的文本作为上下文的一部分,分段继续生成,而不是把输入随机切开。一种通用做法是:每次只生成一定量的 token,然后将其追加到上下文,再取上下文最后若干 token 作为下一次的输入。这个“若干 token”就是要保留的重叠区,用来维持语义连续性。
下面是一个不依赖特定模型 API 的分块生成骨架:
def chunked_generate(model, tokenizer, start_ids, max_total_length, chunk_size, overlap=64):
current_ids = list(start_ids)
while len(current_ids) < max_total_length:
# 只取最后 context_window 个 token,避免上下文无限膨胀
context = current_ids[-(1024 + overlap):]
output_ids = model.generate(
torch.tensor([context]),
max_new_tokens=chunk_size,
do_sample=False
)
new_tokens = output_ids[0].tolist()[len(context):]
current_ids.extend(new_tokens)
if len(new_tokens) == 0:
break
# 下一次使用时会保留重叠区,见上面的 context 切片
return tokenizer.decode(current_ids[start_ids:])这里的 overlap 不是硬性参数,但它要与模型的上下文窗口匹配。如果模型上下文窗口是 2048,而 context_window 设为 1024,那么 overlap 只需要几十到一两百;如果上下文窗口本身很紧,建议至少保留 128 个 token 的重叠区,否则可能丢失前文的关键约束。每次生成后需要检查 len(new_tokens),防止模型提前结束导致死循环。
用生成结果完整性校验替代报错判断
OOM 不报错不代表生成成功。分块策略最怕两件事:模型提前结束、语义断开。提前结束时生成的文本长度会明显小于预期;语义断开则需要检查关键内容是否保留。可以在每个分块后记录输出长度,并在最终结果中检查几个在上下文中出现过的关键片段:
def verify_output(text, expected_keyphrases, min_length):
# 1) 长度检查:至少达到 min_length,建议设为期望长度的较小值
if len(text.split()) < min_length:
return False, f'length too short: {len(text.split())} < {min_length}'
# 2) 关键片段保留检查
missing = [phrase for phrase in expected_keyphrases if phrase not in text]
if missing:
return False, f'missing phrases: {missing}'
return True, 'ok'这段函数需要在生成结束后调用,并且把 min_length 设置为期望 token 数或字符数。关键片段不要选太长,建议选一句话的 5-10 个字;如果分块导致某段丢失,这个检查能快速暴露问题。若失败,先把 overlap 调大,再检查分块之间的拼接逻辑。
记录资源参数并形成启动配置
找到一组能稳定成功的参数后,把它们固化下来,不要每次反复试验。建议写到环境变量和启动脚本里,内容至少包括:max_length、chunk_size、overlap、批量大小(batch_size)以及是否使用 torch.cuda.empty_cache()。例如:
export WE_LM_MAX_LENGTH=512
export WE_LM_CHUNK_SIZE=128
export WE_LM_OVERLAP=64
export WE_LM_BATCH_SIZE=1
export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
python run_welm_generate.py \
`--checkpoint` /path/to/welm \
`--input` prompt.txt \
`--output` output.txt将上述环境变量放在项目根目录的 .env 或启动脚本开头。运行前用 nvidia-smi 确认空闲显存;运行后把生成的日志、csv 文件一起保存。这里的 expandable_segments:True 是为了减少显存碎片,但它让显存分配更激进,不一定适合所有驱动,需要结合环境确认是否启用。