mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 11:16:00 +00:00
feat(train): route reward model training through rewards/factory instead of policies/factory
This commit is contained in:
@@ -1,9 +1,7 @@
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from lerobot.datasets import LeRobotDataset
|
from lerobot.datasets import LeRobotDataset
|
||||||
from lerobot.policies import make_policy
|
from lerobot.rewards import RewardClassifierConfig, make_reward_model, make_reward_pre_post_processors
|
||||||
from lerobot.rewards.classifier.configuration_classifier import RewardClassifierConfig
|
|
||||||
from lerobot.rewards.factory import make_reward_pre_post_processors
|
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
@@ -24,9 +22,9 @@ def main():
|
|||||||
model_name="microsoft/resnet-18",
|
model_name="microsoft/resnet-18",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Make policy, preprocessor, and optimizer
|
# Make reward model, preprocessor, and optimizer
|
||||||
policy = make_policy(config, ds_meta=dataset.meta)
|
reward_model = make_reward_model(config, dataset_stats=dataset.meta.stats)
|
||||||
optimizer = config.get_optimizer_preset().build(policy.parameters())
|
optimizer = config.get_optimizer_preset().build(reward_model.parameters())
|
||||||
preprocessor, _ = make_reward_pre_post_processors(config, dataset_stats=dataset.meta.stats)
|
preprocessor, _ = make_reward_pre_post_processors(config, dataset_stats=dataset.meta.stats)
|
||||||
|
|
||||||
classifier_id = "<user>/reward_classifier_hil_serl_example"
|
classifier_id = "<user>/reward_classifier_hil_serl_example"
|
||||||
@@ -44,7 +42,7 @@ def main():
|
|||||||
batch = preprocessor(batch)
|
batch = preprocessor(batch)
|
||||||
|
|
||||||
# Forward pass
|
# Forward pass
|
||||||
loss, output_dict = policy.forward(batch)
|
loss, output_dict = reward_model.forward(batch)
|
||||||
|
|
||||||
# Backward pass and optimization
|
# Backward pass and optimization
|
||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
@@ -60,8 +58,8 @@ def main():
|
|||||||
|
|
||||||
print("Training finished!")
|
print("Training finished!")
|
||||||
|
|
||||||
# You can now save the trained policy.
|
# You can now save the trained reward model.
|
||||||
policy.push_to_hub(classifier_id)
|
reward_model.push_to_hub(classifier_id)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
@@ -47,7 +47,7 @@ from lerobot.datasets import EpisodeAwareSampler, make_dataset
|
|||||||
from lerobot.envs import close_envs, make_env, make_env_pre_post_processors
|
from lerobot.envs import close_envs, make_env, make_env_pre_post_processors
|
||||||
from lerobot.optim.factory import make_optimizer_and_scheduler
|
from lerobot.optim.factory import make_optimizer_and_scheduler
|
||||||
from lerobot.policies import PreTrainedPolicy, make_policy, make_pre_post_processors
|
from lerobot.policies import PreTrainedPolicy, make_policy, make_pre_post_processors
|
||||||
from lerobot.rewards.factory import make_reward_pre_post_processors
|
from lerobot.rewards import make_reward_pre_post_processors
|
||||||
from lerobot.utils.import_utils import register_third_party_plugins
|
from lerobot.utils.import_utils import register_third_party_plugins
|
||||||
from lerobot.utils.logging_utils import AverageMeter, MetricsTracker
|
from lerobot.utils.logging_utils import AverageMeter, MetricsTracker
|
||||||
from lerobot.utils.random_utils import set_seed
|
from lerobot.utils.random_utils import set_seed
|
||||||
@@ -250,6 +250,17 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
logging.info("Creating env")
|
logging.info("Creating env")
|
||||||
eval_env = make_env(cfg.env, n_envs=cfg.eval.batch_size, use_async_envs=cfg.eval.use_async_envs)
|
eval_env = make_env(cfg.env, n_envs=cfg.eval.batch_size, use_async_envs=cfg.eval.use_async_envs)
|
||||||
|
|
||||||
|
if cfg.is_reward_model_training:
|
||||||
|
if is_main_process:
|
||||||
|
logging.info("Creating reward model")
|
||||||
|
from lerobot.rewards import make_reward_model
|
||||||
|
|
||||||
|
policy = make_reward_model(
|
||||||
|
cfg=cfg.reward_model,
|
||||||
|
dataset_stats=dataset.meta.stats,
|
||||||
|
dataset_meta=dataset.meta,
|
||||||
|
)
|
||||||
|
else:
|
||||||
if is_main_process:
|
if is_main_process:
|
||||||
logging.info("Creating policy")
|
logging.info("Creating policy")
|
||||||
policy = make_policy(
|
policy = make_policy(
|
||||||
@@ -260,16 +271,16 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
|
|
||||||
if cfg.peft is not None:
|
if cfg.peft is not None:
|
||||||
logging.info("Using PEFT! Wrapping model.")
|
logging.info("Using PEFT! Wrapping model.")
|
||||||
# Convert CLI peft config to dict for overrides
|
|
||||||
peft_cli_overrides = dataclasses.asdict(cfg.peft)
|
peft_cli_overrides = dataclasses.asdict(cfg.peft)
|
||||||
policy = policy.wrap_with_peft(peft_cli_overrides=peft_cli_overrides)
|
policy = policy.wrap_with_peft(peft_cli_overrides=peft_cli_overrides)
|
||||||
|
|
||||||
# Wait for all processes to finish policy creation before continuing
|
# Wait for all processes to finish model creation before continuing
|
||||||
accelerator.wait_for_everyone()
|
accelerator.wait_for_everyone()
|
||||||
|
|
||||||
processor_pretrained_path = cfg.policy.pretrained_path
|
active_cfg = cfg.trainable_config
|
||||||
|
processor_pretrained_path = active_cfg.pretrained_path
|
||||||
if (
|
if (
|
||||||
getattr(cfg.policy, "use_relative_actions", False)
|
getattr(active_cfg, "use_relative_actions", False)
|
||||||
and processor_pretrained_path is not None
|
and processor_pretrained_path is not None
|
||||||
and not cfg.resume
|
and not cfg.resume
|
||||||
):
|
):
|
||||||
@@ -279,18 +290,15 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
)
|
)
|
||||||
processor_pretrained_path = None
|
processor_pretrained_path = None
|
||||||
|
|
||||||
# Create processors - only provide dataset_stats if not resuming from saved processors
|
|
||||||
processor_kwargs = {}
|
processor_kwargs = {}
|
||||||
postprocessor_kwargs = {}
|
postprocessor_kwargs = {}
|
||||||
if (processor_pretrained_path and not cfg.resume) or not processor_pretrained_path:
|
if (processor_pretrained_path and not cfg.resume) or not processor_pretrained_path:
|
||||||
# Only provide dataset_stats when not resuming from saved processor state
|
|
||||||
processor_kwargs["dataset_stats"] = dataset.meta.stats
|
processor_kwargs["dataset_stats"] = dataset.meta.stats
|
||||||
|
|
||||||
# For SARM, always provide dataset_meta for progress normalization
|
if cfg.is_reward_model_training and cfg.reward_model.type == "sarm":
|
||||||
if cfg.policy.type == "sarm":
|
|
||||||
processor_kwargs["dataset_meta"] = dataset.meta
|
processor_kwargs["dataset_meta"] = dataset.meta
|
||||||
|
|
||||||
if processor_pretrained_path is not None:
|
if not cfg.is_reward_model_training and processor_pretrained_path is not None:
|
||||||
processor_kwargs["preprocessor_overrides"] = {
|
processor_kwargs["preprocessor_overrides"] = {
|
||||||
"device_processor": {"device": device.type},
|
"device_processor": {"device": device.type},
|
||||||
"normalizer_processor": {
|
"normalizer_processor": {
|
||||||
@@ -310,11 +318,9 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
from lerobot.configs.rewards import RewardModelConfig
|
if cfg.is_reward_model_training:
|
||||||
|
|
||||||
if isinstance(cfg.policy, RewardModelConfig):
|
|
||||||
preprocessor, postprocessor = make_reward_pre_post_processors(
|
preprocessor, postprocessor = make_reward_pre_post_processors(
|
||||||
cfg.policy,
|
cfg.reward_model,
|
||||||
**processor_kwargs,
|
**processor_kwargs,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
Reference in New Issue
Block a user