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

使用jit_compile=True时TFT5ForConditionalGeneration.generate()返回空scores的解决方法

使用jit_compile=True时TFT5ForConditionalGeneration.generate()返回空scores的解决方法

嗨,这个问题我之前也碰到过!本质上是TensorFlow的XLA编译(也就是jit_compile=True)对Hugging Face生成函数返回的动态结构处理有局限——generate返回的scores是一个逐token的logits列表,XLA在编译时没法很好地追踪这种动态长度的嵌套结构,甚至会直接优化掉它认为“未被显式使用”的返回值,导致你拿到空列表。

这里有几个可行的解决办法,你可以根据自己的需求试一下:

方法1:显式提取并转换scores为可追踪的张量格式

在tf.function内部,不要直接返回整个生成结果对象,而是把scores提取出来,转换成TensorFlow能正确识别的张量结构(比如堆叠成一个三维张量),同时返回你需要的其他结果。这样XLA能明确追踪到这个返回值,就不会丢失了:

@tf.function(jit_compile=True)
def generate(transformer_model, input_ids, generation_config):
    generated_output = transformer_model.generate(
        input_ids, generation_config=generation_config, return_dict_in_generate=True, output_scores=True
    )
    # 将scores列表堆叠成一个张量(形状:[生成步数, batch_size, vocab_size])
    stacked_scores = tf.stack(generated_output.scores, axis=0)
    # 返回需要的结果:生成的序列、堆叠后的scores、注意力掩码
    return generated_output.sequences, stacked_scores, generated_output.attention_mask

调用的时候,你可以这样获取结果:

sequences, scores, attn_mask = generate(transformer_model, text_tokenized.input_ids, generation_config)
print(scores)  # 这里就能拿到正常的scores张量了

方法2:让generate部分避开XLA编译

因为generate本身是包含逐token循环的复杂逻辑,XLA对这类动态循环的支持本来就不算完美。你可以只对模型的前向传播部分使用jit_compile,而把generate放在普通Python函数里:

# 只编译模型的前向计算(如果需要的话)
@tf.function(jit_compile=True)
def model_forward(input_ids, attention_mask, labels):
    return transformer_model(input_ids=input_ids, attention_mask=attention_mask, labels=labels)

# generate函数不使用jit_compile
def generate(transformer_model, input_ids, generation_config):
    return transformer_model.generate(
        input_ids, generation_config=generation_config, return_dict_in_generate=True, output_scores=True
    )

这样既保留了前向计算的编译加速,又能正常拿到scores。

方法3:用autograph标记跳过generate的编译

可以在tf.function内部,对transformer.generate的调用标记为不进行XLA转换,让这部分逻辑保持原生执行:

@tf.function(jit_compile=True)
def generate(transformer_model, input_ids, generation_config):
    # 标记内部函数不被autograph转换
    @tf.autograph.experimental.do_not_convert
    def inner_generate():
        return transformer_model.generate(
            input_ids, generation_config=generation_config, return_dict_in_generate=True, output_scores=True
        )
    generated_output = inner_generate()
    return generated_output

这样外层函数的其他部分还是会被XLA编译,但generate部分会按原生逻辑执行,就能正常返回scores了。

额外注意点

  • 确保你的generation_config里静态设置了output_scores=True和return_dict_in_generate=True,不要在tf.function内部动态修改这些参数,XLA对静态参数的处理更稳定。
  • 如果你不需要多生成序列,设置num_return_sequences=1也能减少动态性,帮助XLA更好地处理返回结果。

备注:内容来源于stack exchange,提问作者ayalaall

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.20 13:03:09