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

如何向带Triplet Loss的BERT模型正确输入三元组并训练?

问题分析与解决方案

1. 报错直接原因

Keras的fit()方法强制要求提供目标数据y——哪怕你的损失函数实际用不到y_true。框架默认所有损失函数都接收y_true和y_pred两个参数,所以必须传一个占位符(dummy)值,不能直接传None。

2. 模型核心问题修正

除了dummy y值的问题,你的代码还有两个关键错误需要修复:

2.1 BERT输出未做句子级池化

当前模型直接拼接了BERT输出的序列级embedding(形状(None, 64, 768)),这会导致后续损失计算逻辑完全错误——你需要先将每个句子的序列embedding转换成单个句子向量(比如取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token、全局平均池化或最大池化)。

2.2 VECTOR_SIZE未定义且切分逻辑错误

triplet_loss中使用了未定义的VECTOR_SIZE,而且直接按列切分序列级embedding的方式不符合triplet loss的计算逻辑(应该基于句子级向量计算距离)。

3. 完整修正代码

3.1 修正损失函数与模型结构

import tensorflow as tf
import pandas as pd
from transformers import DistilBertTokenizer, TFDistilBertModel
from tensorflow import keras

model_name = "distilbert-base-multilingual-cased"
tokenizer = DistilBertTokenizer.from_pretrained(model_name)
bert_model = TFDistilBertModel.from_pretrained(model_name)

# 修正数据变量名错误(原代码中test未定义)
test_data = test_data[['anchor', 'match', 'non_match']]
sample_size = int(0.8 * len(test_data))
train, validation = test_data[:sample_size], test_data[sample_size:]

# 定义句子向量池化函数,将序列embedding转为句子级向量
def get_sentence_embedding(bert_output):
    # 方案1:取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token的embedding
    return bert_output[:, 0, :]
    # 方案2:全局平均池化
    # return tf.reduce_mean(bert_output, axis=1)

def triplet_loss(y_true, y_pred):
    """计算Triplet Loss,忽略y_true参数"""
    VECTOR_SIZE = 768  # DistilBERT的隐藏层固定维度
    anchor = y_pred[:, :VECTOR_SIZE]
    positive = y_pred[:, VECTOR_SIZE:2*VECTOR_SIZE]
    negative = y_pred[:, 2*VECTOR_SIZE:]
    
    pos_dist = tf.reduce_sum(tf.square(anchor - positive), axis=1)
    neg_dist = tf.reduce_sum(tf.square(anchor - negative), axis=1)
    
    basic_loss = pos_dist - neg_dist + 0.1  # margin=0.1,可根据需求调整
    loss = tf.reduce_mean(tf.maximum(basic_loss, 0.0))
    return loss

def create_model(bert_model, max_length):
    input_anchor = tf.keras.layers.Input(shape=(max_length,), name='input_anchor', dtype='int32')
    input_positive = tf.keras.layers.Input(shape=(max_length,), name='input_positive', dtype='int32')
    input_negative = tf.keras.layers.Input(shape=(max_length,), name='input_negative', dtype='int32')
    
    # 获取BERT输出并转换为句子向量
    anchor_bert_out = bert_model(input_anchor)[0]
    anchor_emb = get_sentence_embedding(anchor_bert_out)
    
    positive_bert_out = bert_model(input_positive)[0]
    positive_emb = get_sentence_embedding(positive_bert_out)
    
    negative_bert_out = bert_model(input_negative)[0]
    negative_emb = get_sentence_embedding(negative_bert_out)
    
    # 拼接三个句子向量,形状为(None, 3*768)
    merged_output = tf.keras.layers.concatenate([anchor_emb, positive_emb, negative_emb])
    
    model = tf.keras.Model(inputs=[input_anchor, input_positive, input_negative], outputs=merged_output)
    # Triplet Loss场景下accuracy无意义,建议移除该指标
    model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=2e-5), loss=triplet_loss)
    return model

3.2 修正训练代码(补充tokenization步骤)

max_length = 64

# 文本转模型可接收的整数序列
def tokenize_texts(texts):
    return tokenizer(
        texts.tolist(),
        padding='max_length',
        truncation=True,
        max_length=max_length,
        return_tensors='tf'
    )['input_ids']

# 处理训练数据
train_anchor = tokenize_texts(train['anchor'])
train_positive = tokenize_texts(train['match'])
train_negative = tokenize_texts(train['non_match'])

# 创建模型
triplet_model = create_model(bert_model, max_length)
triplet_model.summary()

# 生成dummy y值(全0数组,仅满足框架要求)
dummy_y = tf.zeros((len(train),))

# 开始训练
history = triplet_model.fit(
    x=[train_anchor, train_positive, train_negative],
    y=dummy_y,
    batch_size=32,
    epochs=5,
    # 验证集同样需要传dummy y
    validation_data=([tokenize_texts(validation['anchor']), tokenize_texts(validation['match']), tokenize_texts(validation['non_match'])], tf.zeros((len(validation),)))
)

4. 额外说明

  • 训练时建议冻结BERT的部分层(比如前几层),避免过拟合,可通过bert_model.trainable = False后再解冻部分层实现。
  • 可调整margin值(当前为0.1)和学习率,优化训练效果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 18:04:53