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
相关产品推荐
相关产品推荐

