mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 11:16:00 +00:00
refactor(policies): remove reward model branches from policy factory and __init__
This commit is contained in:
@@ -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",
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user