fix(train): raise ValueError for non-trainable reward models in train function

This commit is contained in:
Khalil Meftah
2026-04-21 14:31:40 +02:00
parent 402f754209
commit c424d906ce
+5
View File
@@ -260,6 +260,11 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
dataset_stats=dataset.meta.stats,
dataset_meta=dataset.meta,
)
if not policy.is_trainable:
raise ValueError(
f"Reward model '{policy.name}' is zero-shot and cannot be trained via lerobot-train. "
"Use it directly for inference via compute_reward() (e.g. offline precompute)."
)
else:
if is_main_process:
logging.info("Creating policy")