[v1] refactor registry plugin structure and params (#10641)

This commit is contained in:
Jiaqi
2026-07-24 15:23:21 +08:00
committed by GitHub
parent 19e9fe3ced
commit 3f77101580
50 changed files with 843 additions and 1007 deletions

View File

@@ -76,8 +76,7 @@ class RMTrainer(BaseTrainer):
train_dataset,
callbacks=None,
) -> None:
cp_size = args.dist_config.get("cp_size", 1) if args.dist_config is not None else 1
if cp_size > 1:
if args.cp_size > 1:
raise NotImplementedError("RM trainer currently only supports cp_size == 1.")
super().__init__(args, model, renderer, train_dataset, callbacks)
@@ -163,7 +162,7 @@ class RMTrainer(BaseTrainer):
def run_rm(args: InputArgument = None):
model_args, data_args, training_args, _ = get_args(args)
model_args.model_class = ModelClass.CLS
DistributedInterface(training_args.dist_config)
DistributedInterface(training_args)
train_dataset = DataEngine(data_args.train_dataset)
_validate_rm_dataset_format(train_dataset, data_args.train_dataset)
model_engine = ModelEngine(model_args, is_train=True)