refactor(policies): remove reward model branches from policy factory and __init__

This commit is contained in:
Khalil Meftah
2026-04-17 15:08:13 +02:00
parent cbaf1acffd
commit 94e20c6ef7
2 changed files with 2 additions and 36 deletions
-7
View File
@@ -12,11 +12,6 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # 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 .act.configuration_act import ACTConfig as ACTConfig
from .diffusion.configuration_diffusion import DiffusionConfig as DiffusionConfig from .diffusion.configuration_diffusion import DiffusionConfig as DiffusionConfig
from .factory import get_policy_class, make_policy, make_policy_config, make_pre_post_processors from .factory import get_policy_class, make_policy, make_policy_config, make_pre_post_processors
@@ -48,9 +43,7 @@ __all__ = [
"PI0Config", "PI0Config",
"PI0FastConfig", "PI0FastConfig",
"PI05Config", "PI05Config",
"RewardClassifierConfig",
"SACConfig", "SACConfig",
"SARMConfig",
"SmolVLAConfig", "SmolVLAConfig",
"TDMPCConfig", "TDMPCConfig",
"VQBeTConfig", "VQBeTConfig",
+2 -29
View File
@@ -36,8 +36,6 @@ from lerobot.processor import (
transition_to_batch, transition_to_batch,
transition_to_policy_action, 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.types import PolicyAction
from lerobot.utils.constants import ( from lerobot.utils.constants import (
ACTION, ACTION,
@@ -89,7 +87,7 @@ def get_policy_class(name: str) -> type[PreTrainedPolicy]:
Args: Args:
name: The name of the policy. Supported names are "tdmpc", "diffusion", "act", 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: Returns:
The policy class corresponding to the given name. 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 from .sac.modeling_sac import SACPolicy
return SACPolicy return SACPolicy
elif name == "reward_classifier":
from lerobot.rewards.classifier.modeling_classifier import Classifier
return Classifier
elif name == "smolvla": elif name == "smolvla":
from .smolvla.modeling_smolvla import SmolVLAPolicy from .smolvla.modeling_smolvla import SmolVLAPolicy
return SmolVLAPolicy return SmolVLAPolicy
elif name == "sarm":
from lerobot.rewards.sarm.modeling_sarm import SARMRewardModel
return SARMRewardModel
elif name == "groot": elif name == "groot":
from .groot.modeling_groot import GrootPolicy from .groot.modeling_groot import GrootPolicy
@@ -173,7 +163,7 @@ def make_policy_config(policy_type: str, **kwargs) -> PreTrainedConfig:
Args: Args:
policy_type: The type of the policy. Supported types include "tdmpc", policy_type: The type of the policy. Supported types include "tdmpc",
"multi_task_dit", "diffusion", "act", "vqbet", "pi0", "pi05", "sac", "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. **kwargs: Keyword arguments to be passed to the configuration class constructor.
Returns: Returns:
@@ -376,14 +366,6 @@ def make_pre_post_processors(
dataset_stats=kwargs.get("dataset_stats"), 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): elif isinstance(policy_cfg, SmolVLAConfig):
from .smolvla.processor_smolvla import make_smolvla_pre_post_processors 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"), 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): elif isinstance(policy_cfg, GrootConfig):
from .groot.processor_groot import make_groot_pre_post_processors from .groot.processor_groot import make_groot_pre_post_processors