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

R Keras中BERT输出与额外特征拼接报错求助

问题原因

你的报错源于管道符%>%的误用:你通过管道将cls_token传入layer_concatenate,同时又指定了inputs参数,这导致函数收到两个冲突的输入——管道传递的未命名参数(即cls_token)和inputs列表,触发Keras的参数校验错误。此外,input_topic是int32类型,而BERT输出为浮点型,拼接前未统一类型也可能引发后续问题。

修正方案

方案1:移除管道符,直接调用拼接层

这是最直接的方式,明确传入所有待拼接的张量列表:

input_word_ids <- layer_input(shape = c(set.max_length), dtype = 'int32', name = "input_word_ids")
input_mask <- layer_input(shape = c(set.max_length), dtype = 'int32', name = "input_attention_mask")
input_topic <- layer_input(shape = c(13), dtype = 'int32', name = "input_topic") %>%
  layer_dtype_cast(dtype = 'float32') # 统一 dtype 避免类型冲突

last_hidden_state <- model_tf(input_word_ids, attention_mask = input_mask)[[1]] # shape=(None, 512, 768)
cls_token <- last_hidden_state[, 1,] # shape=(None, 768)

# 直接调用layer_concatenate,传入张量列表
output <- layer_concatenate(inputs = list(cls_token, input_topic), axis = -1)

方案2:正确使用管道符传递拼接列表

如果坚持用管道,需将待拼接的张量列表作为管道输入,而非单独传递单个张量:

input_word_ids <- layer_input(shape = c(set.max_length), dtype = 'int32', name = "input_word_ids")
input_mask <- layer_input(shape = c(set.max_length), dtype = 'int32', name = "input_attention_mask")
input_topic <- layer_input(shape = c(13), dtype = 'int32', name = "input_topic") %>%
  layer_dtype_cast(dtype = 'float32')

last_hidden_state <- model_tf(input_word_ids, attention_mask = input_mask)[[1]]
cls_token <- last_hidden_state[, 1,]

# 将张量列表传入管道,再调用拼接层
output <- list(cls_token, input_topic) %>% layer_concatenate(axis = -1)
关键说明
  • 必须保证所有待拼接的张量数据类型一致:BERT的输出是float32,因此需要将input_topic通过layer_dtype_cast转为浮点型,否则会出现类型不兼容错误。
  • 管道符%>%的作用是将左侧对象作为右侧函数的第一个参数,如果同时指定inputs参数,就会造成参数重复传入,这是核心错误点。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 14:45:31