You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让自定义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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.21 02:15:57