evo1-meanflow / scripts /train_splitMeanflow.py
SeongJoon's picture
Add Evo1 MeanFlow variants (meanflow / adaptive / nocurr / splitMeanflow)
b1c7880 verified
Raw History Blame Contribute Delete
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)