GR00T N1.7源码学习(一):工程入口、模型结构与动作生成流程解析-CSDN博客

GR00T N1.7源码学习(二):训练数据、Processor与多机器人动作空间解析-CSDN博客

GR00T N1.7源码学习(三):动作头内部模块、DiT结构与多机器人条件编码解析-CSDN博客

前三篇已经分别分析了GR00T N1.7的模型主线、数据处理流程和动作头内部结构。这一篇重点是看一次N1.7微调任务是怎样被真正组装起来的:配置如何进入Pipeline,模型和Processor如何加载,Dataset如何创建,Trainer如何接管训练循环,最后Checkpoint如何保存成后续可以继续训练或部署的目录。对应源码主要集中在,

gr00t/experiment/experiment.py
gr00t/experiment/trainer.py
gr00t/experiment/utils.py
gr00t/experiment/dist_utils.py
gr00t/model/registry.py
gr00t/model/gr00t_n1d7/setup.py
gr00t/data/dataset/factory.py
gr00t/configs/training/training_config.py

1、run(config)串起N1.7微调任务的主流程

launch_finetune.py负责解析命令行参数,并把模型路径、数据集路径、机器人类型、训练哪些模块等信息写入config,最后调用,

run(config)

训练主流程在gr00t/experiment/experiment.py中。函数开头没有立刻创建模型,而是先完成训练前的检查和初始化,

def run(config: Config):
    """Main training function."""
    warn_configs(config)
    check_resume_compatibility(config.training)

    global_rank = _init_distributed_process_group()

    setup_logging()
    if global_rank != 0:
        logging.getLogger().setLevel(logging.WARNING)
    set_seed(config.data.seed)

    config.validate()

这里的warn_configs(config)主要检查训练配置是否合理,例如global_batch_size能否被GPU数量整除、是否开启了梯度累积、视频后端配置是否合适等。check_resume_compatibility(config.training)则检查断点续训配置,例如save_only_model=True不能和严格的resume_from_checkpoint=True一起使用。后面初始化分布式环境、日志系统和随机种子,最后调用config.validate()做配置合法性检查。

warn_configs中有一段和Batch Size相关的检查比较值得注意,

assert config.training.global_batch_size % config.training.num_gpus == 0, (
    "global_batch_size must be divisible by num_gpus"
)

if config.training.gradient_accumulation_steps > 1:
    logging.info(
        "global_batch_size=%d × gradient_accumulation_steps=%d "
        "→ accumulated_batch_size=%d per optimizer step",
        config.training.global_batch_size,
        config.training.gradient_accumulation_steps,
        config.training.accumulated_batch_size,
    )

global_batch_size表示一次forward/backward在所有GPU上的总样本数。如果设置了gradient_accumulation_steps,真实每次参数更新对应的样本数还要再乘以梯度累积步数。这个区别在大模型微调中很常见,尤其是显存放不下大Batch时,通常会把单步Batch设小,再通过梯度累积补等效Batch Size。

2、输出目录保存训练配置和Processor信息

配置检查完成后,run(config)会创建训练输出目录,

if config.training.experiment_name is None:
    output_dir = Path(config.training.output_dir)
    experiment_name = output_dir.name
else:
    output_dir = Path(config.training.output_dir) / config.training.experiment_name
    experiment_name = config.training.experiment_name

output_dir.mkdir(parents=True, exist_ok=True)

save_cfg_dir = output_dir / "experiment_cfg"
processor_dir = output_dir / "processor"

后续模型权重、Processor、实验配置和Checkpoint都会保存在这个目录下面。源码会先把当前训练配置写到experiment_cfg目录中,

run_on_rank0(
    save_run_config_artifacts,
    save_cfg_dir,
    output_dir,
    config,
    experiment_name,
)

save_run_config_artifacts会保存config.yaml、conf.yaml和wandb_config.json,

def save_run_config_artifacts(
    save_cfg_dir: Path, output_dir: Path, config: Config, experiment_name: str
):
    save_cfg_dir.mkdir(parents=True, exist_ok=True)
    config.save(save_cfg_dir / "config.yaml")
    omegaconf_config = OmegaConf.create(config.__dict__)
    omegaconf_config["max_steps"] = config.training.max_steps
    omegaconf_config["save_steps"] = config.training.save_steps
    OmegaConf.save(omegaconf_config, save_cfg_dir / "conf.yaml", resolve=True)

    wandb_config_file = output_dir / "wandb_config.json"
    with open(wandb_config_file, "w") as f:
        json.dump(
            {
                "project": config.training.wandb_project,
                "run_id": experiment_name,
            },
            f,
        )

3、MODEL_REGISTRY选择Gr00tN1d7Pipeline

训练对象开始组装的位置是下面几行,

pipeline = MODEL_REGISTRY.get(type(config.model))(config, save_cfg_dir)
pipeline.setup()
model = pipeline.return_model()
train_dataset, eval_dataset = pipeline.return_dataset()
data_collator = pipeline.return_collator()
processor = pipeline.return_processor()

通过MODEL_REGISTRY根据模型配置选择对应Pipeline。注册表定义在gr00t/model/registry.py:

MODEL_REGISTRY = {}

def register_model(model_cfg_cls, pipeline_cls):
    if model_cfg_cls in MODEL_REGISTRY:
        raise ValueError(f"Model type '{model_cfg_cls}' already registered.")
    MODEL_REGISTRY[model_cfg_cls] = pipeline_cls

N1.7对应的注册发生在gr00t/model/gr00t_n1d7/setup.py中,

register_model(Gr00tN1d7Config, Gr00tN1d7Pipeline)

因此当type(config.model)是Gr00tN1d7Config时,训练流程会自动选择Gr00tN1d7Pipeline。这层设计把通用训练流程和具体模型实现解耦了。后续如果新增其他模型,只要注册新的Config -> Pipeline映射,experiment.run()主流程可以基本保持不变。

Gr00tN1d7Pipeline的结构比如下,

class Gr00tN1d7Pipeline(ModelPipeline):
    model_class = Gr00tN1d7
    processor_class = Gr00tN1d7Processor

    def __init__(self, config: Config, save_cfg_dir: Path):
        super().__init__(config)
        self.save_cfg_dir = save_cfg_dir

setup()主要完成三件事,

def setup(self):
    self.model = self._create_model()
    self.train_dataset, self.eval_dataset = self._create_dataset(self.save_cfg_dir)
    self.data_collator = self._create_collator()

4、Pipeline加载模型、Processor和训练数据集

Gr00tN1d7Pipeline._create_model()负责创建模型。如果配置了start_from_checkpoint并且没有跳过权重加载,源码会从Checkpoint恢复模型,

model, loading_info = AutoModel.from_pretrained(
    self.config.training.start_from_checkpoint,
    tune_llm=self.config.model.tune_llm,
    tune_visual=self.config.model.tune_visual,
    tune_projector=self.config.model.tune_projector,
    tune_diffusion_model=self.config.model.tune_diffusion_model,
    tune_vlln=self.config.model.tune_vlln,
    state_dropout_prob=self.config.model.state_dropout_prob,
    backbone_trainable_params_fp32=self.config.model.backbone_trainable_params_fp32,
    load_bf16=self.config.model.load_bf16,
    transformers_loading_kwargs=self.transformers_loading_kwargs,
    output_loading_info=True,
    **self.transformers_loading_kwargs,
)

模型加载完成后,源码会检查Checkpoint权重是否和当前代码匹配,

missing_keys = loading_info.get("missing_keys", [])
unexpected_keys = loading_info.get("unexpected_keys", [])
mismatched_keys = loading_info.get("mismatched_keys", [])
other_missing = [k for k in missing_keys if "mask_token" not in k]

errors = []
if other_missing:
    errors.append(f"Missing keys ({len(other_missing)}): {other_missing}")
if unexpected_keys:
    errors.append(f"Unexpected keys ({len(unexpected_keys)}): {unexpected_keys}")
if mismatched_keys:
    errors.append(f"Mismatched keys ({len(mismatched_keys)}): {mismatched_keys}")
if errors:
    raise RuntimeError(
        "Checkpoint weight mismatch for "
        f"{self.config.training.start_from_checkpoint}:\n" + "\n".join(errors)
    )

除了mask_token缺失会做特殊兼容,其他权重缺失、意外权重或shape不匹配都会直接报错。这种处理比较稳,尤其是GR00T代码版本更新较快时,Checkpoint和当前源码不匹配最好尽早暴露。

Processor的加载方式和模型类似。如果从Checkpoint继续训练,会先从Checkpoint恢复Processor,再用当前训练配置覆盖数据模态、图像增强、动作horizon、相对动作等参数,

processor = AutoProcessor.from_pretrained(
    self.config.training.start_from_checkpoint,
    modality_configs=self.config.data.modality_configs,
    use_percentiles=self.model_config.use_percentiles,
    image_crop_size=self.model_config.image_crop_size,
    image_target_size=self.model_config.image_target_size,
    random_rotation_angle=self.model_config.random_rotation_angle,
    color_jitter_params=self.model_config.color_jitter_params,
    model_name=self.model_config.model_name,
    model_type=self.model_config.backbone_model_type,
    formalize_language=self.model_config.formalize_language,
    max_action_horizon=self.model_config.action_horizon,
    use_relative_action=self.model_config.use_relative_action,
    exclude_state=self.model_config.exclude_state,
    state_dropout_prob=self.model_config.state_dropout_prob,
    use_mean_std=self.model_config.use_mean_std,
    **self.transformers_loading_kwargs,
)

训练时Processor不只是一个临时预处理函数,它会随着Checkpoint一起保存,部署时还要继续负责状态归一化、动作反归一化、相对动作还原和模态字段解析。

数据集通过DatasetFactory构造,

dataset_factory = DatasetFactory(config=self.config)
train_dataset, eval_dataset = dataset_factory.build(processor=self.processor)

DatasetFactory会先生成普通统计量和相对动作统计量,

with run_or_wait_on_rank0(label=f"generate_stats({dataset_path})") as is_rank0:
    if is_rank0:
        generate_stats(dataset_path)
        generate_rel_stats(dataset_path, EmbodimentTag(embodiment_tag))

然后为每个数据集路径创建ShardedSingleStepDataset,

dataset = ShardedSingleStepDataset(
    dataset_path=dataset_path,
    embodiment_tag=EmbodimentTag(embodiment_tag),
    modality_configs=self.config.data.modality_configs[embodiment_tag],
    video_backend=self.config.data.video_backend,
    shard_size=self.config.data.shard_size,
    episode_sampling_rate=self.config.data.episode_sampling_rate,
    seed=self.config.data.seed,
    allow_padding=self.config.data.allow_padding,
)

如果有多个数据集,会根据长度和mix_ratio计算权重,最后组合成ShardedMixtureDataset,

ShardedMixtureDataset(
    datasets=all_datasets,
    weights=all_weights,
    processor=processor,
    seed=self.config.data.seed,
    training=True,
    num_shards_per_epoch=self.config.data.num_shards_per_epoch,
    override_pretraining_statistics=self.config.data.override_pretraining_statistics,
)

5、TrainingArguments承接优化器、精度和分布式参数

模型、Processor、Dataset和Collator准备好以后,run(config)会构造Hugging Face的TrainingArguments。在这之前,源码先处理DeepSpeed配置和每卡Batch Size,

if config.training.num_gpus > 1 and not config.training.use_ddp:
    deepspeed_config = config.get_deepspeed_config()
else:
    deepspeed_config = None

if config.training.per_gpu_batch_size is None:
    per_device_train_batch_size = config.training.global_batch_size // config.training.num_gpus
else:
    per_device_train_batch_size = config.training.per_gpu_batch_size

随后创建TrainingArguments,

training_args = TrainingArguments(
    output_dir=str(output_dir),
    max_steps=config.training.max_steps,
    per_device_train_batch_size=per_device_train_batch_size,
    gradient_accumulation_steps=config.training.gradient_accumulation_steps,
    learning_rate=config.training.learning_rate,
    lr_scheduler_type=config.training.lr_scheduler_type,
    weight_decay=config.training.weight_decay,
    warmup_ratio=config.training.warmup_ratio,
    max_grad_norm=config.training.max_grad_norm,
    logging_steps=config.training.logging_steps,
    save_steps=config.training.save_steps,
    save_total_limit=config.training.save_total_limit,
    save_only_model=config.training.save_only_model,
    fp16=config.training.fp16,
    bf16=config.training.bf16,
    tf32=config.training.tf32,
    gradient_checkpointing=config.training.gradient_checkpointing,
    optim=config.training.optim,
    dataloader_num_workers=config.training.dataloader_num_workers,
    report_to="wandb" if config.training.use_wandb else "none",
    seed=config.data.seed,
    deepspeed=deepspeed_config,
    ddp_find_unused_parameters=False,
    ddp_bucket_cap_mb=config.training.ddp_bucket_cap_mb,
    eval_strategy=config.training.eval_strategy,
    remove_unused_columns=config.training.remove_unused_columns,
    ignore_data_skip=True,
)

6、Gr00tTrainer重写DataLoader创建和断点续训逻辑

训练器创建如下,

trainer = Gr00tTrainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset,
    eval_dataset=eval_dataset,
    data_collator=data_collator,
    multiprocessing_context=config.data.multiprocessing_context,
)

Gr00tTrainer继承自Hugging Face Trainer,但重写了DataLoader创建逻辑,

def get_train_dataloader(self):
    """Return a iterable dataloader without skipping the data during resume, but reseed the dataset instead."""

    self.args.ignore_data_skip = True
    curr_global_step = self.state.global_step
    print(f"Current global step: {curr_global_step}")
    if curr_global_step > 0:
        new_seed = self.train_dataset.seed + curr_global_step
        self.train_dataset.reset_seed(new_seed)
        print(
            f"Resetting seed to {new_seed}. Please note that this will make the experiment non-reproducible."
        )

    data_collator = self.data_collator
    data_collator = self._get_collator_with_removed_columns(
        data_collator, description="training"
    )
    persistent_workers = self.args.dataloader_num_workers > 0

    dataloader_params = {
        "batch_size": self._train_batch_size,
        "collate_fn": data_collator,
        "num_workers": self.args.dataloader_num_workers,
        "pin_memory": self.args.dataloader_pin_memory,
        "persistent_workers": persistent_workers,
    }

    if self.args.dataloader_num_workers > 0:
        dataloader_params["multiprocessing_context"] = self.multiprocessing_context

    return torch.utils.data.DataLoader(self.train_dataset, **dataloader_params)

恢复训练时,它不会让Trainer按默认逻辑跳过已经训练过的数据,而是根据global_step重置数据集seed。源码注释中特别强调,不同rank必须使用相同的新seed,否则ShardedMixtureDataset的分片计划会不一致,可能导致样本重复或丢失。

Gr00tTrainer.train()也重写了一部分逻辑,主要是为了在创建DataLoader之前先加载TrainerState,

def train(
    self,
    resume_from_checkpoint=None,
    **kwargs,
):
    if resume_from_checkpoint is True:
        latest_checkpoint = get_last_checkpoint(self.args.output_dir)
        if latest_checkpoint is None:
            raise ValueError(
                f"No valid checkpoint found in output directory ({self.args.output_dir})"
            )
    elif resume_from_checkpoint in (False, None):
        latest_checkpoint = None
    else:
        latest_checkpoint = resume_from_checkpoint

    if latest_checkpoint is not None:
        logging.info(f"Resuming from checkpoint {latest_checkpoint}")
        self.state = TrainerState.load_from_json(
            os.path.join(latest_checkpoint, TRAINER_STATE_NAME)
        )

    return super().train(resume_from_checkpoint=latest_checkpoint, **kwargs)

get_train_dataloader()读取self.state.global_step时,拿到的是恢复后的真实步数。这里还有一个训练配置检查,

def check_resume_compatibility(training: TrainingConfig) -> None:
    if training.save_only_model and training.resume_from_checkpoint:
        raise ValueError(
            "save_only_model=True is incompatible with resume_from_checkpoint=True..."
        )

原因是save_only_model=True只保存模型权重,不保存optimizer、scheduler和随机数状态,不能用于严格断点续训。

训练Loss本身仍然由模型forward()返回,Trainer只是接管通用反向传播流程,

def compute_loss(
    self,
    model,
    inputs,
    return_outputs: bool = False,
    num_items_in_batch: int | None = None,
):
    loss, outputs = super().compute_loss(
        model,
        inputs,
        return_outputs=True,
        num_items_in_batch=num_items_in_batch,
    )

    self.loss = loss

7、Checkpoint同时保存模型权重、Processor和实验配置

创建Trainer之后,源码会添加一个保存格式回调,

trainer.add_callback(
    CheckpointFormatCallback(
        run_name=experiment_name,
        exp_cfg_dir=save_cfg_dir,
        processor_dir=processor_dir,
    )
)

CheckpointFormatCallback定义在gr00t/experiment/utils.py,作用是在每次保存Checkpoint后,把实验配置和Processor复制进去,

class CheckpointFormatCallback(TrainerCallback):
    """This callback format checkpoint to make them standalone."""

    def on_save(self, args, state, control, **kwargs):
        if state.is_world_process_zero:
            checkpoint_dir = Path(args.output_dir) / f"checkpoint-{state.global_step}"

复制配置目录,

if self.exp_cfg_dir is not None:
    exp_cfg_dst = checkpoint_dir / self.exp_cfg_dir.name
    if self.exp_cfg_dir.exists():
        shutil.copytree(self.exp_cfg_dir, exp_cfg_dst, dirs_exist_ok=True)

复制Processor目录,

if self.processor_dir is not None:
    if self.processor_dir.exists():
        shutil.copytree(self.processor_dir, checkpoint_dir, dirs_exist_ok=True)

复制W&B配置,

wandb_config_src = Path(args.output_dir) / "wandb_config.json"
wandb_config_dst = checkpoint_dir / "wandb_config.json"
if wandb_config_src.exists():
    shutil.copy2(wandb_config_src, wandb_config_dst)

部署时并不是只需要模型权重,还需要Processor、统计量、模态配置、动作维度和相对动作配置。如果Checkpoint目录里只有权重文件,后面Policy推理可能能加载模型,但动作归一化、反归一化或字段解析会出问题。

训练开始前,源码也会先保存一次Processor,

run_on_rank0(processor.save_pretrained, processor_dir, label="processor.save_pretrained")

源码注释里提到,Processor中的statistics.json会被Gr00tPolicy.from_pretrained读取。如果这个文件写坏或缺失,推理动作可能不会立刻报错,但机器人执行出来的动作会不对。

训练结束后,源码保存最终模型,并关闭数据集资源,

trainer.save_model()
logging.info(f"Model saved to {output_dir}")

if hasattr(train_dataset, "close"):
    train_dataset.close()
if eval_dataset is not None and hasattr(eval_dataset, "close"):
    eval_dataset.close()
logging.info("Training completed!")
Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐