如何让自定义collate_fn适配不同的填充(pad)ID?
让collate_fn适配不同填充ID的方法
有两种简单实用的方式可以让你的collate_fn灵活适配不同模型的填充ID:
方法一:通过闭包传入填充ID
写一个生成collate_fn的外层函数,将填充ID作为参数传入,内层函数直接使用该参数完成填充逻辑:
def create_collate_fn(pad_id): def collate_fn(batch): max_len = max([len(b['input_ids']) for b in batch]) # 修正原代码的循环问题,遍历整个batch处理每个样本 input_ids = [b['input_ids'] + ([pad_id] * (max_len - len(b['input_ids']))) for b in batch] labels = [b['label'] for b in batch] return {'input_ids': input_ids, 'labels': labels} return collate_fn # 使用示例:根据目标模型的分词器获取填充ID from transformers import AutoTokenizer tokenizer = AutoTokenizer.from_pretrained("bert-base-chinese") # 生成对应填充ID的collate_fn target_collate_fn = create_collate_fn(tokenizer.pad_token_id) # 传入DataLoader使用 from torch.utils.data import DataLoader dataloader = DataLoader(your_dataset, batch_size=8, collate_fn=target_collate_fn)
方法二:直接传入分词器到collate_fn
如果不想额外写外层函数,也可以直接把分词器作为参数传入collate_fn,利用分词器自带的pad_token_id属性:
def collate_fn(batch, tokenizer): max_len = max([len(b['input_ids']) for b in batch]) pad_id = tokenizer.pad_token_id input_ids = [b['input_ids'] + ([pad_id] * (max_len - len(b['input_ids']))) for b in batch] labels = [b['label'] for b in batch] return {'input_ids': input_ids, 'labels': labels} # 使用示例:用lambda包装传入分词器 tokenizer = AutoTokenizer.from_pretrained("gpt2") dataloader = DataLoader(your_dataset, batch_size=8, collate_fn=lambda x: collate_fn(x, tokenizer))
注:原示例代码中input_ids的列表推导缺少遍历逻辑,已在上述代码中修正,确保处理batch内所有样本。
内容的提问来源于stack exchange,提问作者Sean
相关产品推荐
相关产品推荐

