跑通配置和训练脚本
This commit is contained in:
@@ -1,45 +1,108 @@
|
||||
import hydra
|
||||
from omegaconf import DictConfig, OmegaConf
|
||||
from hydra.utils import instantiate
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
import logging
|
||||
import hydra
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
from omegaconf import DictConfig, OmegaConf
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.optim import AdamW
|
||||
|
||||
# 必须指向你的配置文件所在路径
|
||||
# config_path 是相对于当前脚本的路径,或者绝对路径
|
||||
# config_name 是不带 .yaml 后缀的主文件名
|
||||
@hydra.main(version_base=None, config_path="../../roboimi/vla/conf", config_name="config")
|
||||
# 确保导入路径正确
|
||||
sys.path.append(os.getcwd())
|
||||
|
||||
from roboimi.vla.agent import VLAAgent
|
||||
from hydra.utils import instantiate
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@hydra.main(version_base=None, config_path="../../../roboimi/vla/conf", config_name="config")
|
||||
def main(cfg: DictConfig):
|
||||
print(f"Working directory : {os.getcwd()}")
|
||||
print(f"Configuration:\n{OmegaConf.to_yaml(cfg)}")
|
||||
print(OmegaConf.to_yaml(cfg))
|
||||
log.info(f"🚀 Starting VLA Training with Real Data (Device: {cfg.train.device})")
|
||||
|
||||
# 1. 实例化 Agent
|
||||
# Hydra 会自动查找 _target_ 并递归实例化 vlm_backbone 和 action_head
|
||||
print(">>> Instantiating VLA Agent...")
|
||||
agent = instantiate(cfg.agent)
|
||||
# --- 1. 实例化 Dataset & DataLoader ---
|
||||
# Hydra 根据 conf/data/custom_hdf5.yaml 实例化类
|
||||
dataset = instantiate(cfg.data)
|
||||
|
||||
# 将模型移至 GPU
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
agent.to(device)
|
||||
print(f">>> Agent created successfully. Backbone: {type(agent.vlm).__name__}")
|
||||
|
||||
# 2. 实例化 DataLoader (假设你也为 Data 写了 yaml)
|
||||
# 实例化 Dataset
|
||||
dataset = hydra.utils.instantiate(cfg.data)
|
||||
|
||||
# 封装进 DataLoader
|
||||
dataloader = torch.utils.data.DataLoader(
|
||||
dataset,
|
||||
batch_size=cfg.train.batch_size,
|
||||
dataloader = DataLoader(
|
||||
dataset,
|
||||
batch_size=cfg.train.batch_size,
|
||||
shuffle=True,
|
||||
num_workers=4
|
||||
num_workers=cfg.train.num_workers,
|
||||
pin_memory=(cfg.train.device != "cpu")
|
||||
)
|
||||
log.info(f"✅ Dataset loaded. Size: {len(dataset)}")
|
||||
|
||||
# 3. 实例化 Optimizer (Hydra 也支持 partial 实例化)
|
||||
# optimizer = instantiate(cfg.train.optimizer, params=agent.parameters())
|
||||
# --- 2. 实例化 Agent ---
|
||||
agent: VLAAgent = instantiate(cfg.agent)
|
||||
agent.to(cfg.train.device)
|
||||
agent.train()
|
||||
|
||||
# 4. 模拟训练循环
|
||||
print(f">>> Starting training with batch size: {cfg.train.batch_size}")
|
||||
# ... training loop logic here ...
|
||||
optimizer = AdamW(agent.parameters(), lr=cfg.train.lr)
|
||||
|
||||
# --- 3. Training Loop ---
|
||||
# 使用一个无限迭代器或者 epoch 循环
|
||||
data_iter = iter(dataloader)
|
||||
pbar = tqdm(range(cfg.train.max_steps), desc="Training")
|
||||
|
||||
for step in pbar:
|
||||
try:
|
||||
batch = next(data_iter)
|
||||
except StopIteration:
|
||||
#而在 epoch 结束时重新开始
|
||||
data_iter = iter(dataloader)
|
||||
batch = next(data_iter)
|
||||
|
||||
# Move to device
|
||||
# 注意:这里需要递归地将字典里的 tensor 移到 GPU
|
||||
batch = recursive_to_device(batch, cfg.train.device)
|
||||
|
||||
# --- 4. Adapter Layer (适配层) ---
|
||||
# Dataset 返回的是具体的相机 key (如 'agentview_image' 或 'top')
|
||||
# Agent 期望的是通用的 'image'
|
||||
# 我们在这里做一个映射,模拟多模态融合前的处理
|
||||
|
||||
# 假设我们只用配置里的第一个 key 作为主视觉
|
||||
primary_cam_key = cfg.data.obs_keys[0]
|
||||
|
||||
# Dataset 返回 shape: (B, Obs_Horizon, C, H, W)
|
||||
# DebugBackbone 期望: (B, C, H, W) 或者 (B, Seq, Dim)
|
||||
# 这里我们取 Obs_Horizon 的最后一帧 (Current Frame)
|
||||
input_img = batch['obs'][primary_cam_key][:, -1, :, :, :]
|
||||
|
||||
agent_input = {
|
||||
"obs": {
|
||||
"image": input_img,
|
||||
"text": batch["language"] # 传递语言指令
|
||||
},
|
||||
"actions": batch["actions"] # (B, Chunk, Dim)
|
||||
}
|
||||
|
||||
# --- 5. Forward & Backward ---
|
||||
outputs = agent(agent_input)
|
||||
|
||||
# 处理 Loss 掩码 (如果在真实训练中,需要在这里应用 action_mask)
|
||||
# 目前 DebugHead 内部直接算了 MSE,还没用 mask,我们在下一阶段优化 Policy 时加上
|
||||
loss = outputs['loss']
|
||||
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
if step % cfg.train.log_freq == 0:
|
||||
pbar.set_postfix({"loss": f"{loss.item():.4f}"})
|
||||
|
||||
log.info("✅ Training Loop with Real HDF5 Finished!")
|
||||
|
||||
def recursive_to_device(data, device):
|
||||
if isinstance(data, torch.Tensor):
|
||||
return data.to(device)
|
||||
elif isinstance(data, dict):
|
||||
return {k: recursive_to_device(v, device) for k, v in data.items()}
|
||||
elif isinstance(data, list):
|
||||
return [recursive_to_device(v, device) for v in data]
|
||||
return data
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user