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

如何在R语言的TensorFlow Dataset中配置样本权重

正确配置带样本权重的TensorFlow数据集(R语言)

问题根源

你之前错误地将sample_weight归入输出字典,导致Keras将其识别为模型的额外输出节点,与模型实际输出output_tensor不匹配,从而触发报错。当使用Dataset作为fit输入时,样本权重需要作为数据集元素的第三个独立组成部分,结构应为(输入字典, 标签字典, 样本权重张量)。

修正后的完整代码

1. 构造带权重的训练/测试数据集

with(tf$device("/cpu:0"),
  train_dataset <- tensor_slices_dataset(
    list(
      list("input_tensor" = image_data[!test_set,,,]),  # 输入特征字典
      list("output_tensor" = label_data[!test_set,]),    # 标签字典
      sample_weights[!test_set]                          # 单独的样本权重张量
    )
  ) %>% 
  dataset_shuffle(1024) %>% 
  dataset_batch(batch_size)
)

with(tf$device("/cpu:0"),
  test_dataset <- tensor_slices_dataset(
    list(
      list("input_tensor" = image_data[test_set,,,]),
      list("output_tensor" = label_data[test_set,]),
      sample_weights[test_set]
    )
  ) %>% 
  dataset_shuffle(1024) %>% 
  dataset_batch(batch_size)
)

2. 验证数据集结构

修正后train_dataset的element_spec会显示为:

<BatchDataset element_spec=({'input_tensor': TensorSpec(shape=(None, 50, 50, 3), dtype=tf.float64, name=None)}, 
                             {'output_tensor': TensorSpec(shape=(None, 4), dtype=tf.float64, name=None)},
                             TensorSpec(shape=(None,), dtype=tf.float64, name=None))>

3. 正常执行模型训练

原有的fit代码无需修改,直接运行即可:

history <- model %>% fit(
  train_dataset,
  epochs = number_epochs,
  view_metrics = TRUE,
  verbose = 1,
  validation_data = test_dataset
)

之前尝试失败的原因解析

  • 错误嵌入样本权重到输出字典:Keras会将输出字典的所有键视为模型输出节点,而你的模型仅定义了output_tensor一个输出,因此触发不匹配错误。
  • from_tensor_slices用法错误:你传入的嵌套列表结构不符合要求,正确结构应为list(输入, 标签, 权重),而非多层嵌套的list。
  • 直接在fit中传sample_weights参数:当输入为Dataset时,Keras不支持单独传入该参数,必须将权重嵌入到Dataset的元素结构中。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 05:44:51