Download scripts/train_splitMeanflow.py from SeongJoon/evo1-meanflow: direct link, hf CLI and curl.
- Browser
- Download file 20.8 kB
-
https://huggingface.co/SeongJoon/evo1-meanflow/resolve/main/scripts/train_splitMeanflow.py
- Command line
-
hf download hf://SeongJoon/evo1-meanflow/scripts/train_splitMeanflow.py
-
curl -L -o train_splitMeanflow.py https://huggingface.co/SeongJoon/evo1-meanflow/resolve/main/scripts/train_splitMeanflow.py
20.8 kB
| """ | |
| SplitMeanFlow νμ΅ μ€ν¬λ¦½νΈ | |
| train_meanflow.pyμμ μ°¨μ΄μ : | |
| 1. EVO1_SplitMeanFlow λͺ¨λΈ μ¬μ© | |
| 2. compute_splitMeanflow_loss() νΈμΆ (JVP μμ) | |
| 3. JVP μμΌλ―λ‘ bfloat16 autocast μ μ§ κ°λ₯ | |
| (μμ μ±μ μν΄ float32 μ μ§) | |
| """ | |
| import sys | |
| import os | |
| import math | |
| sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))) | |
| import torch | |
| import torch.nn as nn | |
| from torch.utils.data import DataLoader | |
| from tqdm import tqdm | |
| from torch.optim.lr_scheduler import LambdaLR | |
| from torch.optim import AdamW | |
| from accelerate import Accelerator, DistributedType | |
| import logging | |
| import argparse | |
| import json | |
| import shutil | |
| import warnings | |
| from Evo1_splitMeanflow import EVO1_SplitMeanFlow | |
| accelerator = Accelerator() | |
| def get_with_warning(config, key, default): | |
| if key in config: | |
| return config[key] | |
| warnings.warn(f"'{key}' not found in config, using default: {default!r}") | |
| return default | |
| def custom_collate_fn(batch): | |
| prompts = [item["prompt"] for item in batch] | |
| images = [item["images"] for item in batch] | |
| states = torch.stack([item["state"] for item in batch]) | |
| actions = torch.stack([item["action"] for item in batch]) | |
| action_mask = torch.stack([item["action_mask"] for item in batch]) | |
| image_masks = torch.stack([item["image_mask"] for item in batch]) | |
| state_mask = torch.stack([item["state_mask"] for item in batch]) | |
| embodiment_ids = torch.stack([item["embodiment_id"] for item in batch]) | |
| return { | |
| "prompts": prompts, "images": images, | |
| "states": states, "actions": actions, | |
| "action_mask": action_mask, "state_mask": state_mask, | |
| "image_masks": image_masks, "embodiment_ids": embodiment_ids, | |
| } | |
| def get_lr_lambda(warmup_steps, total_steps, resume_step=0): | |
| def lr_lambda(current_step): | |
| current_step += resume_step | |
| if current_step < warmup_steps: | |
| return current_step / max(1, warmup_steps) | |
| progress = (current_step - warmup_steps) / max(1, total_steps - warmup_steps) | |
| return max(0.0, 0.5 * (1.0 + math.cos(math.pi * progress))) | |
| return lr_lambda | |
| def setup_logging(log_dir): | |
| from datetime import datetime | |
| timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") | |
| os.makedirs(log_dir, exist_ok=True) | |
| log_path = os.path.join(log_dir, f"train_smf_{timestamp}.log") | |
| if accelerator.is_main_process: | |
| logging.basicConfig( | |
| level=logging.INFO, | |
| format="%(asctime)s [%(levelname)s] %(message)s", | |
| handlers=[logging.FileHandler(log_path), logging.StreamHandler()], | |
| ) | |
| logging.info(f"[SplitMeanFlow] Logging to: {log_path}") | |
| return log_path | |
| def prepare_dataset(config): | |
| dataset_type = get_with_warning(config, "dataset_type", "lerobot") | |
| image_size = get_with_warning(config, "image_size", 448) | |
| max_samples = get_with_warning(config, "max_samples_per_file", None) | |
| horizon = get_with_warning(config, "horizon", 50) | |
| binarize_gripper = get_with_warning(config, "binarize_gripper", False) | |
| use_augmentation = get_with_warning(config, "use_augmentation", False) | |
| if dataset_type == "lerobot": | |
| from dataset.lerobot_dataset_pretrain_mp import LeRobotDataset | |
| import yaml | |
| with open(config.get("dataset_config_path"), "r") as f: | |
| dataset_config = yaml.safe_load(f) | |
| dataset = LeRobotDataset( | |
| config=dataset_config, | |
| image_size=image_size, | |
| max_samples_per_file=max_samples, | |
| action_horizon=horizon, | |
| binarize_gripper=binarize_gripper, | |
| use_augmentation=use_augmentation, | |
| ) | |
| else: | |
| raise ValueError(f"Unknown dataset_type: {dataset_type}") | |
| if accelerator.is_main_process: | |
| logging.info(f"Loaded {len(dataset)} samples ({dataset_type})") | |
| return dataset | |
| def prepare_dataloader(dataset, config): | |
| batch_size = get_with_warning(config, "batch_size", 8) | |
| num_workers = get_with_warning(config, "num_workers", 8) | |
| dataloader = DataLoader( | |
| dataset, | |
| batch_size=batch_size, | |
| shuffle=True, | |
| num_workers=num_workers, | |
| pin_memory=True, | |
| persistent_workers=False, | |
| drop_last=True, | |
| collate_fn=custom_collate_fn, | |
| ) | |
| if accelerator.is_main_process: | |
| logging.info(f"Dataloader: batch_size={batch_size}") | |
| return dataloader | |
| def check_numerical_stability(step, **named_tensors): | |
| for name, tensor in named_tensors.items(): | |
| if not torch.isfinite(tensor).all(): | |
| logging.warning(f"[Step {step}] Non-finite in {name}") | |
| return False | |
| return True | |
| def build_param_groups(model, wd): | |
| decay, no_decay = [], [] | |
| for n, p in model.named_parameters(): | |
| if not p.requires_grad: | |
| continue | |
| if n.endswith("bias") or "norm" in n.lower() or p.dim() == 1: | |
| no_decay.append(p) | |
| else: | |
| decay.append(p) | |
| return [{"params": decay, "weight_decay": wd}, | |
| {"params": no_decay, "weight_decay": 0.0}] | |
| def get_and_clip_grad_norm(accelerator, model, loss, max_norm=1.0): | |
| grad_norms = [p.grad.norm(2) for p in model.parameters() if p.grad is not None] | |
| if not grad_norms: | |
| total_norm = clipped_norm = torch.tensor(0.0, device=loss.device) | |
| else: | |
| total_norm = torch.norm(torch.stack(grad_norms), 2) | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm) | |
| clipped_norm = torch.norm( | |
| torch.stack([p.grad.norm(2) for p in model.parameters() if p.grad is not None]), 2 | |
| ) | |
| return total_norm, clipped_norm | |
| def save_checkpoint(save_dir, step, model_engine, loss, accelerator, config=None, norm_stats=None): | |
| tag = f"step_{step}" | |
| checkpoint_dir = os.path.join(save_dir, tag) | |
| if accelerator.is_main_process and os.path.exists(checkpoint_dir): | |
| shutil.rmtree(checkpoint_dir) | |
| accelerator.wait_for_everyone() | |
| client_state = { | |
| "step": step, | |
| "best_loss": loss if isinstance(loss, float) else loss.item(), | |
| "config": config, | |
| } if accelerator.is_main_process else {} | |
| if hasattr(model_engine, "save_checkpoint"): | |
| model_engine.save_checkpoint(save_dir, tag=tag, client_state=client_state) | |
| else: | |
| unwrapped = accelerator.unwrap_model(model_engine) | |
| accelerator.save_model(unwrapped, os.path.join(save_dir, tag)) | |
| if accelerator.is_main_process: | |
| torch.save(client_state, os.path.join(checkpoint_dir, "training_state.pt")) | |
| if accelerator.is_main_process: | |
| if config is not None: | |
| with open(os.path.join(checkpoint_dir, "config.json"), "w") as f: | |
| json.dump(config, f, indent=2) | |
| if norm_stats is not None: | |
| with open(os.path.join(checkpoint_dir, "norm_stats.json"), "w") as f: | |
| json.dump(norm_stats, f, indent=2) | |
| with open(os.path.join(checkpoint_dir, "checkpoint.json"), "w") as f: | |
| json.dump({"type": "ds_model", "version": 0.0, | |
| "checkpoints": "mp_rank_00_model_states.pt"}, f, indent=2) | |
| logging.info(f"Saved checkpoint β {checkpoint_dir}") | |
| if isinstance(step, int) or (isinstance(step, str) and step.startswith("epoch_") and step[6:].isdigit()): | |
| import pathlib | |
| root = pathlib.Path(save_dir) | |
| numeric_ckpts = [] | |
| for d in root.iterdir(): | |
| if not d.is_dir(): | |
| continue | |
| for prefix in ("step_", "epoch_"): | |
| if d.name.startswith(prefix): | |
| suffix = d.name[len(prefix):] | |
| try: | |
| numeric_ckpts.append((int(suffix), d)) | |
| except ValueError: | |
| pass | |
| numeric_ckpts.sort(key=lambda x: x[0], reverse=True) | |
| for _, old_dir in numeric_ckpts[4:]: | |
| shutil.rmtree(old_dir) | |
| logging.info(f"Removed old checkpoint: {old_dir}") | |
| def load_checkpoint_with_deepspeed(model_engine, load_dir, accelerator, | |
| tag="step_best", load_optimizer_states=True, | |
| resume_pretrain=False): | |
| try: | |
| load_path, client_state = model_engine.load_checkpoint( | |
| load_dir, tag=tag, load_module_strict=True, | |
| load_optimizer_states=load_optimizer_states and not resume_pretrain, | |
| load_lr_scheduler_states=load_optimizer_states and not resume_pretrain, | |
| ) | |
| if accelerator.is_main_process: | |
| logging.info(f"Loaded checkpoint: {load_dir}/{tag}") | |
| return client_state.get("step", 0), client_state | |
| except Exception as e: | |
| if accelerator.is_main_process: | |
| logging.warning(f"Checkpoint load with optimizer failed: {e}. Retrying model-only...") | |
| load_path, client_state = model_engine.load_checkpoint( | |
| load_dir, tag=tag, load_module_strict=True, | |
| load_optimizer_states=False, load_lr_scheduler_states=False, | |
| ) | |
| return client_state.get("step", 0), client_state | |
| # ββ νμ΅ λ©μΈ βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def train(config): | |
| save_dir = get_with_warning(config, "save_dir", "checkpoints_smf") | |
| setup_logging(save_dir) | |
| if get_with_warning(config, "debug", False): | |
| torch.autograd.set_detect_anomaly(True) | |
| dataset = prepare_dataset(config) | |
| dataloader = prepare_dataloader(dataset, config) | |
| model = EVO1_SplitMeanFlow(config) | |
| model.train() | |
| model.set_finetune_flags() | |
| lr = get_with_warning(config, "lr", 1e-5) | |
| wd = get_with_warning(config, "weight_decay", 1e-5) | |
| optimizer = AdamW(build_param_groups(model, wd), lr=lr) | |
| if accelerator.is_main_process: | |
| logging.info(f"[SplitMeanFlow] AdamW lr={lr}, wd={wd}") | |
| model, optimizer, dataloader = accelerator.prepare(model, optimizer, dataloader) | |
| model_engine = model | |
| max_epochs = config.get("max_epochs", None) | |
| if max_epochs is not None: | |
| steps_per_epoch = len(dataloader) | |
| max_steps = int(math.ceil(max_epochs * steps_per_epoch)) | |
| if accelerator.is_main_process: | |
| logging.info(f"{max_epochs} epochs Γ {steps_per_epoch} steps/epoch = {max_steps} total") | |
| else: | |
| max_steps = get_with_warning(config, "max_steps", 1000) | |
| warmup_steps = get_with_warning(config, "warmup_steps", 300) | |
| log_interval = get_with_warning(config, "log_interval", 100) | |
| ckpt_interval = get_with_warning(config, "ckpt_interval", 1000) | |
| ckpt_epoch_interval = get_with_warning(config, "ckpt_epoch_interval", None) | |
| ckpt_epoch_start = get_with_warning(config, "ckpt_epoch_start", 0.0) | |
| max_norm = get_with_warning(config, "grad_clip_norm", 1.0) | |
| if ckpt_epoch_interval is not None: | |
| steps_per_epoch_for_ckpt = len(dataloader) | |
| ckpt_interval = int(ckpt_epoch_interval * steps_per_epoch_for_ckpt) | |
| ckpt_start_step = int(ckpt_epoch_start * steps_per_epoch_for_ckpt) | |
| if accelerator.is_main_process: | |
| logging.info( | |
| f"[Checkpoint] epoch {ckpt_epoch_start}λΆν° " | |
| f"λ§€ {ckpt_epoch_interval} epoch = {ckpt_interval} steps" | |
| ) | |
| else: | |
| ckpt_start_step = 0 | |
| os.makedirs(save_dir, exist_ok=True) | |
| best_loss = float("inf") | |
| resume = get_with_warning(config, "resume", False) | |
| resume_path = get_with_warning(config, "resume_path", None) | |
| resume_pretrain = get_with_warning(config, "resume_pretrain", False) | |
| if resume != bool(resume_path): | |
| raise ValueError("--resume and --resume_path must be set together.") | |
| if resume: | |
| resume_path = resume_path.rstrip("/") | |
| resume_dir, resume_tag = os.path.split(resume_path) | |
| step, client_state = load_checkpoint_with_deepspeed( | |
| model_engine, resume_dir, accelerator, resume_tag, | |
| load_optimizer_states=True, resume_pretrain=resume_pretrain, | |
| ) | |
| best_loss = client_state.get("best_loss", float("inf")) | |
| if accelerator.is_main_process: | |
| logging.info(f"Resumed from {resume_path}, step={step}") | |
| else: | |
| step = 0 | |
| if accelerator.is_main_process: | |
| logging.info("Starting fresh SplitMeanFlow training") | |
| if resume_pretrain: | |
| step = 0 | |
| scheduler = LambdaLR(optimizer, get_lr_lambda(warmup_steps, max_steps, resume_step=step)) | |
| run_name = config.get("run_name", "SplitMeanFlow") | |
| steps_per_epoch = len(dataloader) | |
| pbar = tqdm( | |
| total=max_steps, initial=step, desc=run_name, | |
| disable=not accelerator.is_main_process, | |
| dynamic_ncols=True, smoothing=0.1, | |
| ) | |
| unwrapped_model = accelerator.unwrap_model(model_engine) | |
| while step < max_steps: | |
| for batch in dataloader: | |
| if step >= max_steps: | |
| break | |
| prompts = batch["prompts"] | |
| images_batch = batch["images"] | |
| image_masks = batch["image_masks"] | |
| states = batch["states"] | |
| actions_gt = batch["actions"] | |
| action_mask = batch["action_mask"] | |
| # ββ VLM μΈμ½λ© (bfloat16) ββββββββββββββββββββββββββββββββββββββββ | |
| fused_tokens_list = [] | |
| with torch.amp.autocast(device_type="cuda", dtype=torch.bfloat16): | |
| for prompt, images, image_mask in zip(prompts, images_batch, image_masks): | |
| fused = unwrapped_model.get_vl_embeddings( | |
| images=images, image_mask=image_mask, prompt=prompt, | |
| return_cls_only=False, | |
| ) | |
| fused_tokens_list.append(fused) | |
| # float32 μ μΊμ€νΈ (νμ΅ μμ μ±) | |
| fused_tokens = torch.cat(fused_tokens_list, dim=0).float() | |
| states = states.float() | |
| actions_gt = actions_gt.float() | |
| # ββ SplitMeanFlow μμ€ (JVP μμ, float32) βββββββββββββββββββββββ | |
| loss, loss_dict = unwrapped_model.compute_splitMeanflow_loss( | |
| fused_tokens = fused_tokens, | |
| states = states, | |
| actions_gt = actions_gt, | |
| action_mask = action_mask, | |
| step = step, | |
| total_steps = max_steps, | |
| ) | |
| if not check_numerical_stability(step, loss=loss): | |
| logging.warning(f"[Step {step}] Skipping due to non-finite loss") | |
| continue | |
| optimizer.zero_grad(set_to_none=True) | |
| accelerator.backward(loss) | |
| total_norm, _ = get_and_clip_grad_norm(accelerator, model, loss, max_norm) | |
| optimizer.step() | |
| scheduler.step() | |
| loss_val = loss.item() | |
| if step % log_interval == 0 and accelerator.is_main_process: | |
| epoch = step / steps_per_epoch | |
| logging.info( | |
| f"[Step {step}] Loss: {loss_val:.4f} | " | |
| f"smf_loss={loss_dict['smf_loss']:.4f} | " | |
| f"max_gap={loss_dict['max_gap']:.3f} | " | |
| f"mean_gap={loss_dict['mean_gap']:.3f} | " | |
| f"fm_ratio={loss_dict['fm_ratio']:.2f} | " | |
| f"ep_ratio={loss_dict['ep_ratio']:.2f} | " | |
| f"epoch={epoch:.2f} | " | |
| f"lr={scheduler.get_last_lr()[0]:.2e} | " | |
| f"grad_norm={total_norm:.3f}" | |
| ) | |
| if accelerator.is_main_process: | |
| is_best = loss_val < best_loss | |
| if is_best: | |
| best_loss = loss_val | |
| is_best_tensor = torch.tensor(int(is_best), device=accelerator.device) | |
| else: | |
| is_best_tensor = torch.tensor(0, device=accelerator.device) | |
| if accelerator.distributed_type != DistributedType.NO: | |
| torch.distributed.broadcast(is_best_tensor, src=0) | |
| if is_best_tensor.item() == 1 and step > 1000: | |
| save_checkpoint( | |
| save_dir, step="best", model_engine=model_engine, | |
| loss=loss, accelerator=accelerator, config=config, | |
| norm_stats=dataset.arm2stats_dict, | |
| ) | |
| if accelerator.is_main_process: | |
| logging.info(f"Saved best checkpoint at step {step}, loss={loss_val:.6f}") | |
| step += 1 | |
| pbar.update(1) | |
| pbar.set_postfix( | |
| loss=f"{loss_val:.4f}", | |
| gap=f"{loss_dict['max_gap']:.2f}", | |
| lr=f"{scheduler.get_last_lr()[0]:.2e}", | |
| epoch=f"{step / steps_per_epoch:.2f}", | |
| ) | |
| if step % ckpt_interval == 0 and step > 0 and step >= ckpt_start_step: | |
| if ckpt_epoch_interval is not None: | |
| current_epoch = int(step / steps_per_epoch) | |
| ckpt_tag = f"epoch_{current_epoch}" | |
| else: | |
| ckpt_tag = step | |
| save_checkpoint( | |
| save_dir, step=ckpt_tag, model_engine=model_engine, | |
| loss=loss, accelerator=accelerator, config=config, | |
| norm_stats=dataset.arm2stats_dict, | |
| ) | |
| if accelerator.is_main_process: | |
| epoch_disp = step / steps_per_epoch | |
| logging.info(f"Periodic checkpoint: tag={ckpt_tag} (epoch={epoch_disp:.1f})") | |
| pbar.close() | |
| save_checkpoint( | |
| save_dir, step="final", model_engine=model_engine, | |
| loss=loss, accelerator=accelerator, config=config, | |
| norm_stats=dataset.arm2stats_dict, | |
| ) | |
| if accelerator.is_main_process: | |
| logging.info(f"Training done. best_loss={best_loss:.6f}") | |
| # ββ argparse ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| if __name__ == "__main__": | |
| parser = argparse.ArgumentParser(description="Train EVO1 with SplitMeanFlow") | |
| default_cfg = os.path.abspath( | |
| os.path.join(os.path.dirname(__file__), "..", "dataset", "config.yaml") | |
| ) | |
| parser.add_argument("--device", type=str, default="cuda") | |
| parser.add_argument("--run_name", type=str, default="Evo1_splitMeanflow") | |
| parser.add_argument("--vlm_name", type=str, default="OpenGVLab/InternVL3-1B") | |
| parser.add_argument("--dataset_config_path", type=str, default=default_cfg) | |
| parser.add_argument("--finetune_vlm", type=lambda x: x.lower() == "true", default=False) | |
| parser.add_argument("--finetune_action_head", type=lambda x: x.lower() == "true", default=True) | |
| parser.add_argument("--state_dim", type=int, default=8) | |
| parser.add_argument("--batch_size", type=int, default=4) | |
| parser.add_argument("--max_epochs", type=int, default=None) | |
| parser.add_argument("--max_steps", type=int, default=None) | |
| parser.add_argument("--warmup_steps", type=int, default=1000) | |
| parser.add_argument("--log_interval", type=int, default=10) | |
| parser.add_argument("--ckpt_interval", type=int, default=25000) | |
| parser.add_argument("--ckpt_epoch_interval", type=float,default=None) | |
| parser.add_argument("--ckpt_epoch_start", type=float,default=0.0) | |
| parser.add_argument("--save_dir", type=str, default="checkpoints_smf") | |
| parser.add_argument("--lr", type=float,default=1e-5) | |
| parser.add_argument("--weight_decay", type=float,default=1e-5) | |
| parser.add_argument("--grad_clip_norm", type=float,default=1.0) | |
| parser.add_argument("--use_augmentation", action="store_true") | |
| parser.add_argument("--resume", action="store_true") | |
| parser.add_argument("--resume_pretrain", action="store_true") | |
| parser.add_argument("--resume_path", type=str, default=None) | |
| parser.add_argument("--cfg_scale", type=float,default=2.0) | |
| args = parser.parse_args() | |
| config = vars(args) | |
| config["dataset_type"] = "lerobot" | |
| train(config) | |