GR00T N1.7源码学习(四):微调流程、训练Pipeline与Checkpoint保存机制解析
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!")更多推荐
所有评论(0)