feat(vla): vla框架初始化

This commit is contained in:
gouhanke
2026-02-03 14:18:30 +08:00
parent c1ce560b32
commit 57acfd645f
40 changed files with 443 additions and 63 deletions
+45
View File
@@ -0,0 +1,45 @@
import hydra
from omegaconf import DictConfig, OmegaConf
from hydra.utils import instantiate
import torch
import os
# 必须指向你的配置文件所在路径
# config_path 是相对于当前脚本的路径,或者绝对路径
# config_name 是不带 .yaml 后缀的主文件名
@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)}")
# 1. 实例化 Agent
# Hydra 会自动查找 _target_ 并递归实例化 vlm_backbone 和 action_head
print(">>> Instantiating VLA Agent...")
agent = instantiate(cfg.agent)
# 将模型移至 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,
shuffle=True,
num_workers=4
)
# 3. 实例化 Optimizer (Hydra 也支持 partial 实例化)
# optimizer = instantiate(cfg.train.optimizer, params=agent.parameters())
# 4. 模拟训练循环
print(f">>> Starting training with batch size: {cfg.train.batch_size}")
# ... training loop logic here ...
if __name__ == "__main__":
main()