mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 11:16:00 +00:00
refactor(rewards): expand __init__ facade and fix SARMConfig __post_init__ crash
This commit is contained in:
@@ -13,9 +13,24 @@
|
|||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
from .classifier.configuration_classifier import RewardClassifierConfig as RewardClassifierConfig
|
from .classifier.configuration_classifier import RewardClassifierConfig as RewardClassifierConfig
|
||||||
|
from .factory import (
|
||||||
|
get_reward_model_class as get_reward_model_class,
|
||||||
|
make_reward_model as make_reward_model,
|
||||||
|
make_reward_model_config as make_reward_model_config,
|
||||||
|
make_reward_pre_post_processors as make_reward_pre_post_processors,
|
||||||
|
)
|
||||||
|
from .pretrained import PreTrainedRewardModel as PreTrainedRewardModel
|
||||||
from .sarm.configuration_sarm import SARMConfig as SARMConfig
|
from .sarm.configuration_sarm import SARMConfig as SARMConfig
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
# Configuration classes
|
||||||
"RewardClassifierConfig",
|
"RewardClassifierConfig",
|
||||||
"SARMConfig",
|
"SARMConfig",
|
||||||
|
# Base class
|
||||||
|
"PreTrainedRewardModel",
|
||||||
|
# Factory functions
|
||||||
|
"get_reward_model_class",
|
||||||
|
"make_reward_model",
|
||||||
|
"make_reward_model_config",
|
||||||
|
"make_reward_pre_post_processors",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -108,8 +108,6 @@ class SARMConfig(RewardModelConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
super().__post_init__()
|
|
||||||
|
|
||||||
if self.annotation_mode not in ["single_stage", "dense_only", "dual"]:
|
if self.annotation_mode not in ["single_stage", "dense_only", "dual"]:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"annotation_mode must be 'single_stage', 'dense_only', or 'dual', got {self.annotation_mode}"
|
f"annotation_mode must be 'single_stage', 'dense_only', or 'dual', got {self.annotation_mode}"
|
||||||
|
|||||||
Reference in New Issue
Block a user