feat(vla): vla框架初始化
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user