TensorFlow中map调用encode_batch报错:缺少3个必需参数问题
问题解决:RetinaNet encode_batch 调用参数错误
错误原因
你在调用train_dataset.map(label_encoder.encode_batch())时,错误地执行了函数并将返回值传给了map方法。map需要接收的是一个可调用的函数对象,而非函数执行后的结果,这才导致了参数缺失的报错。
正确写法
情况1:参数顺序完全匹配
如果你的train_dataset每个元素的结构正好是(batch_images, gt_boxes, cls_ids),且与encode_batch的参数顺序一致,直接传递函数引用即可(不要加括号):
train_dataset = train_dataset.map(label_encoder.encode_batch)
情况2:参数需显式解包
如果需要明确指定参数对应关系,或者dataset元素顺序与函数参数不匹配,用lambda解包元素:
train_dataset = train_dataset.map(lambda imgs, boxes, ids: label_encoder.encode_batch(imgs, boxes, ids))
额外验证步骤
确保padded_batch后的dataset结构正确,每个元素确实包含三个对应输入:
# 取出一个元素验证结构 for batch in train_dataset.take(1): print(f"元素数量:{len(batch)}") # 应输出3 print(f"图像形状:{batch[0].shape}") print(f"框形状:{batch[1].shape}") print(f"类别ID形状:{batch[2].shape}")
内容的提问来源于stack exchange,提问作者ZueLukas
相关产品推荐
相关产品推荐

