如何为gensim.models.FastText.train()添加训练进度条?
解决FastText训练大数据集的进度追踪问题
方法一:使用gensim内置的回调接口(推荐)
gensim的FastText继承自Word2Vec体系,提供了CallbackAny2Vec类来自定义训练过程中的进度监控,这是最官方的解决方式,代码实现如下:
from gensim.models.callbacks import CallbackAny2Vec import time # 自定义进度回调类 class TrainProgressCallback(CallbackAny2Vec): def __init__(self, total_epochs, total_examples): self.epoch = 0 self.total_epochs = total_epochs self.total_examples = total_examples self.epoch_start_time = None def on_epoch_begin(self, model): self.epoch_start_time = time.time() print(f"开始训练第 {self.epoch + 1}/{self.total_epochs} 轮") def on_epoch_end(self, model): epoch_duration = time.time() - self.epoch_start_time print(f"第 {self.epoch + 1}/{self.total_epochs} 轮训练完成,耗时 {epoch_duration:.2f} 秒") self.epoch += 1 def on_batch_end(self, model): # 每完成一个批次就输出当前进度(可选,频繁输出可能影响性能) processed = model.corpus_count - model.trainables.remaining_examples progress = (processed / self.total_examples) * 100 print(f"\r当前批次进度: {progress:.2f}%", end="") # 初始化模型 embed_model = FastText(vector_size=meta_hyper['vector_size'], window=meta_hyper['window'], alpha= meta_hyper['alpha'], workers=meta_hyper['CPU']) embed_model.build_vocab(data) # 初始化回调实例 callback = TrainProgressCallback(total_epochs=meta_hyper['epochs'], total_examples=len(data)) # 训练时传入回调 start = time.time() embed_model.train(data, total_examples=len(data), epochs=meta_hyper['epochs'], callbacks=[callback]) total_duration = time.time() - start print(f"\n全部训练完成,总耗时 {total_duration:.2f} 秒")
这个方法能精准追踪每一轮epoch的开始/结束时间,还可以选择监控每个批次的进度。如果觉得批次级的输出太频繁,可以注释掉on_batch_end里的打印逻辑,只保留epoch级的进度。
方法二:手动拆分训练循环,用tqdm包装
如果不想用回调,也可以手动分epoch循环训练,每一轮用tqdm来包装数据集迭代器,直观看到单轮的进度:
from tqdm import tqdm embed_model = FastText(vector_size=meta_hyper['vector_size'], window=meta_hyper['window'], alpha= meta_hyper['alpha'], workers=meta_hyper['CPU']) embed_model.build_vocab(data) start = time.time() alpha = meta_hyper['alpha'] for epoch in range(meta_hyper['epochs']): print(f"第 {epoch+1}/{meta_hyper['epochs']} 轮训练") # 用tqdm包装数据集,显示单轮进度 embed_model.train(tqdm(data, total=len(data)), total_examples=len(data), epochs=1, start_alpha=alpha, end_alpha=alpha) # 固定alpha避免自动衰减,手动控制的话可以调整 alpha *= 0.9 # 可选:手动模拟alpha衰减,和默认逻辑一致 total_duration = time.time() - start print(f"\n全部训练完成,总耗时 {total_duration:.2f} 秒")
注意:这种方式需要手动处理学习率alpha的衰减(默认FastText每轮会自动降低alpha),如果要和原训练逻辑完全一致,需要手动模拟这个衰减过程,或者在每轮训练时设置start_alpha和end_alpha为当前的alpha值。
为什么直接用tqdm包装data没用?
因为gensim的train方法内部会对数据集做二次处理(比如按workers拆分、打乱顺序),直接把data传给tqdm后,gensim实际读取的是tqdm包装后的迭代器,但内部的处理逻辑会导致进度条无法正确统计实际处理的样本量,所以才会失效。
内容的提问来源于stack exchange,提问作者The-One-Who-Speaks-And-Depicts
相关产品推荐
相关产品推荐

