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

TensorFlow下微调HuggingFace遥感DistilBERT预训练模型报错求助

报错核心原因

两个代码逻辑问题直接触发运行错误:

  • HuggingFace 提供的TensorFlow版本预训练模型,必须接收遵循(输入特征字典, 标签)元组格式的tf.data.Dataset输入。你当前构建数据集时把input_ids、attention_mask、labels三个对象平铺传入,模型无法正确识别输入字段映射,直接触发入参不匹配报错。
  • 预处理代码存在两处隐患:一是for循环下的分词逻辑没有正确缩进,实际运行会先触发语法错误;二是使用了已在新版Transformers库中废弃的pad_to_max_length参数,且未配置超长文本截断逻辑,容易出现张量维度对齐异常。
可直接运行的修复方案

1. 修正数据预处理逻辑

替换废弃参数、补全必要配置、修正张量拼接逻辑,代码如下:

input_ids_t = []
attention_masks_t = []

# 注意Python对缩进敏感,for循环内的逻辑必须缩进
for sent in df_train['text_a']:
    encoded_dict = tokenizer.encode_plus(
                        sent,                      
                        add_special_tokens = True, 
                        max_length = 128,           
                        padding = 'max_length', # 替换已废弃的pad_to_max_length参数
                        truncation = True, # 新增超长文本截断,避免长度超过128的文本触发维度报错
                        return_attention_mask = True,  
                        return_tensors = 'tf',     
                   )
    # 去掉单条文本编码后多余的batch维度,避免后续拼接维度异常
    input_ids_t.append(tf.squeeze(encoded_dict['input_ids']))
    attention_masks_t.append(tf.squeeze(encoded_dict['attention_mask']))

# 堆叠单条样本的向量得到整体特征张量
input_ids_t = tf.stack(input_ids_t, axis=0)
attention_masks_t = tf.stack(attention_masks_t, axis=0)
labels_t = np.asarray(df_train['label'])

测试集预处理逻辑和上述代码保持一致即可。

2. 修正TF数据集构建格式

按照模型要求的输入结构封装特征,建议同时补全训练必要的shuffle、batch配置:

# 训练集:封装为(特征字典, 标签)格式
train_data = tf.data.Dataset.from_tensor_slices((
    {
        "input_ids": input_ids_t,
        "attention_mask": attention_masks_t
    },
    labels_t
)).shuffle(buffer_size=len(labels_t)).batch(batch_size=16)

# 测试集不需要shuffle,其余结构一致
test_data = tf.data.Dataset.from_tensor_slices((
    {
        "input_ids": input_ids_test,
        "attention_mask": attention_masks_test
    },
    labels_test
)).batch(batch_size=16)

3. 训练调用注意事项

  • 加载TFDistilBertForSequenceClassification时,需要将num_labels参数设置为你自有数据集的实际标签类别数
  • 调用model.fit()时直接传入构建好的train_data,通过validation_data参数传入test_data即可,无需额外拆分x、y入参

排错提示:如果修改后仍有维度相关报错,先打印三个张量的shape做校验:单标签分类场景下input_ids_t和attention_masks_t的维度应为(样本总数, 128),labels_t维度应为(样本总数,);如果是多标签分类,需要对应调整模型损失函数和标签维度配置。

内容的提问来源于stack exchange,提问作者Fahad Alghamdi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.02 23:31:03