mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 11:16:00 +00:00
refactor(imports): update reward model imports to new module structure
This commit is contained in:
@@ -36,6 +36,8 @@ 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,
|
||||||
@@ -52,8 +54,6 @@ from .pi0.configuration_pi0 import PI0Config
|
|||||||
from .pi05.configuration_pi05 import PI05Config
|
from .pi05.configuration_pi05 import PI05Config
|
||||||
from .pretrained import PreTrainedPolicy
|
from .pretrained import PreTrainedPolicy
|
||||||
from .sac.configuration_sac import SACConfig
|
from .sac.configuration_sac import SACConfig
|
||||||
from .sac.reward_model.configuration_classifier import RewardClassifierConfig
|
|
||||||
from .sarm.configuration_sarm import SARMConfig
|
|
||||||
from .smolvla.configuration_smolvla import SmolVLAConfig
|
from .smolvla.configuration_smolvla import SmolVLAConfig
|
||||||
from .tdmpc.configuration_tdmpc import TDMPCConfig
|
from .tdmpc.configuration_tdmpc import TDMPCConfig
|
||||||
from .utils import validate_visual_features_consistency
|
from .utils import validate_visual_features_consistency
|
||||||
@@ -133,7 +133,7 @@ def get_policy_class(name: str) -> type[PreTrainedPolicy]:
|
|||||||
|
|
||||||
return SACPolicy
|
return SACPolicy
|
||||||
elif name == "reward_classifier":
|
elif name == "reward_classifier":
|
||||||
from .sac.reward_model.modeling_classifier import Classifier
|
from lerobot.rewards.classifier.modeling_classifier import Classifier
|
||||||
|
|
||||||
return Classifier
|
return Classifier
|
||||||
elif name == "smolvla":
|
elif name == "smolvla":
|
||||||
@@ -141,7 +141,7 @@ def get_policy_class(name: str) -> type[PreTrainedPolicy]:
|
|||||||
|
|
||||||
return SmolVLAPolicy
|
return SmolVLAPolicy
|
||||||
elif name == "sarm":
|
elif name == "sarm":
|
||||||
from .sarm.modeling_sarm import SARMRewardModel
|
from lerobot.rewards.sarm.modeling_sarm import SARMRewardModel
|
||||||
|
|
||||||
return SARMRewardModel
|
return SARMRewardModel
|
||||||
elif name == "groot":
|
elif name == "groot":
|
||||||
@@ -377,7 +377,7 @@ def make_pre_post_processors(
|
|||||||
)
|
)
|
||||||
|
|
||||||
elif isinstance(policy_cfg, RewardClassifierConfig):
|
elif isinstance(policy_cfg, RewardClassifierConfig):
|
||||||
from .sac.reward_model.processor_classifier import make_classifier_processor
|
from lerobot.rewards.classifier.processor_classifier import make_classifier_processor
|
||||||
|
|
||||||
processors = make_classifier_processor(
|
processors = make_classifier_processor(
|
||||||
config=policy_cfg,
|
config=policy_cfg,
|
||||||
@@ -393,7 +393,7 @@ def make_pre_post_processors(
|
|||||||
)
|
)
|
||||||
|
|
||||||
elif isinstance(policy_cfg, SARMConfig):
|
elif isinstance(policy_cfg, SARMConfig):
|
||||||
from .sarm.processor_sarm import make_sarm_pre_post_processors
|
from lerobot.rewards.sarm.processor_sarm import make_sarm_pre_post_processors
|
||||||
|
|
||||||
processors = make_sarm_pre_post_processors(
|
processors = make_sarm_pre_post_processors(
|
||||||
config=policy_cfg,
|
config=policy_cfg,
|
||||||
|
|||||||
Reference in New Issue
Block a user