""" 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)