From 94e20c6ef77c70827b306f74b94d21749c9fc766 Mon Sep 17 00:00:00 2001 From: Khalil Meftah Date: Fri, 17 Apr 2026 15:08:13 +0200 Subject: [PATCH] refactor(policies): remove reward model branches from policy factory and __init__ --- src/lerobot/policies/__init__.py | 7 ------- src/lerobot/policies/factory.py | 31 ++----------------------------- 2 files changed, 2 insertions(+), 36 deletions(-) diff --git a/src/lerobot/policies/__init__.py b/src/lerobot/policies/__init__.py index 38c4a4979..d048286dd 100644 --- a/src/lerobot/policies/__init__.py +++ b/src/lerobot/policies/__init__.py @@ -12,11 +12,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -from lerobot.rewards.classifier.configuration_classifier import ( - RewardClassifierConfig as RewardClassifierConfig, -) -from lerobot.rewards.sarm.configuration_sarm import SARMConfig as SARMConfig - from .act.configuration_act import ACTConfig as ACTConfig from .diffusion.configuration_diffusion import DiffusionConfig as DiffusionConfig from .factory import get_policy_class, make_policy, make_policy_config, make_pre_post_processors @@ -48,9 +43,7 @@ __all__ = [ "PI0Config", "PI0FastConfig", "PI05Config", - "RewardClassifierConfig", "SACConfig", - "SARMConfig", "SmolVLAConfig", "TDMPCConfig", "VQBeTConfig", diff --git a/src/lerobot/policies/factory.py b/src/lerobot/policies/factory.py index 79262816b..5be3bca43 100644 --- a/src/lerobot/policies/factory.py +++ b/src/lerobot/policies/factory.py @@ -36,8 +36,6 @@ from lerobot.processor import ( transition_to_batch, transition_to_policy_action, ) -from lerobot.rewards.classifier.configuration_classifier import RewardClassifierConfig -from lerobot.rewards.sarm.configuration_sarm import SARMConfig from lerobot.types import PolicyAction from lerobot.utils.constants import ( ACTION, @@ -89,7 +87,7 @@ def get_policy_class(name: str) -> type[PreTrainedPolicy]: Args: name: The name of the policy. Supported names are "tdmpc", "diffusion", "act", - "multi_task_dit", "vqbet", "pi0", "pi05", "sac", "reward_classifier", "smolvla", "wall_x". + "multi_task_dit", "vqbet", "pi0", "pi05", "sac", "smolvla", "wall_x". Returns: The policy class corresponding to the given name. @@ -132,18 +130,10 @@ def get_policy_class(name: str) -> type[PreTrainedPolicy]: from .sac.modeling_sac import SACPolicy return SACPolicy - elif name == "reward_classifier": - from lerobot.rewards.classifier.modeling_classifier import Classifier - - return Classifier elif name == "smolvla": from .smolvla.modeling_smolvla import SmolVLAPolicy return SmolVLAPolicy - elif name == "sarm": - from lerobot.rewards.sarm.modeling_sarm import SARMRewardModel - - return SARMRewardModel elif name == "groot": from .groot.modeling_groot import GrootPolicy @@ -173,7 +163,7 @@ def make_policy_config(policy_type: str, **kwargs) -> PreTrainedConfig: Args: policy_type: The type of the policy. Supported types include "tdmpc", "multi_task_dit", "diffusion", "act", "vqbet", "pi0", "pi05", "sac", - "smolvla", "reward_classifier", "wall_x". + "smolvla", "wall_x". **kwargs: Keyword arguments to be passed to the configuration class constructor. Returns: @@ -376,14 +366,6 @@ def make_pre_post_processors( dataset_stats=kwargs.get("dataset_stats"), ) - elif isinstance(policy_cfg, RewardClassifierConfig): - from lerobot.rewards.classifier.processor_classifier import make_classifier_processor - - processors = make_classifier_processor( - config=policy_cfg, - dataset_stats=kwargs.get("dataset_stats"), - ) - elif isinstance(policy_cfg, SmolVLAConfig): from .smolvla.processor_smolvla import make_smolvla_pre_post_processors @@ -392,15 +374,6 @@ def make_pre_post_processors( dataset_stats=kwargs.get("dataset_stats"), ) - elif isinstance(policy_cfg, SARMConfig): - from lerobot.rewards.sarm.processor_sarm import make_sarm_pre_post_processors - - processors = make_sarm_pre_post_processors( - config=policy_cfg, - dataset_stats=kwargs.get("dataset_stats"), - dataset_meta=kwargs.get("dataset_meta"), - ) - elif isinstance(policy_cfg, GrootConfig): from .groot.processor_groot import make_groot_pre_post_processors