LingArm_frozen_preview / config.yaml
LingArm's picture
Upload folder using huggingface_hub
372c993 verified
Raw
History Blame Contribute Delete
2.35 kB
# Frozen VLA Training Configuration (Qwen3-VL Version)
# PAI Environment: PyTorch 2.6.0 + Python 3.10 + CUDA 12.6 + A10 24GB
#
# 注意:
# - Qwen3-VL-2B-Instruct 已下载到 /mnt/workspace/Qwen3-VL-2B-Instruct/
# - 使用本地路径加载,避免训练时联网下载
# - BridgeData 在 /mnt/workspace/Dataset/stage_frozen01/full/
# Data paths
data_dir: /mnt/workspace/Dataset/stage_frozen01/full
output_dir: /mnt/workspace/checkpoints
log_dir: /mnt/workspace/logs
# Model config
model:
llm: "/mnt/workspace/Qwen3-VL-2B-Instruct" # PAI 本地已下载路径
mlp_hidden_dim: 512
mlp_depth: 2
action_dim: 7 # EEF delta: [dx,dy,dz,ax,ay,az,gripper]
use_processor: true # 使用模型自带 AutoProcessor 处理图像
# Training config
training:
batch_size: 32
gradient_accumulation_steps: 2 # 等效 batch_size=64,省显存
num_epochs: 5
learning_rate: 1.0e-4
weight_decay: 1.0e-2
warmup_steps: 500
max_grad_norm: 1.0
save_every_n_epochs: 1
save_every_n_steps: 500 # 每 500 步额外存一次(防崩,只保留最近3个)
eval_every_n_epochs: 1
# Optimizer
optimizer: "adamw"
scheduler: "cosine"
# Memory optimization
mixed_precision: "bf16" # bf16 比 fp16 更稳定,A10 支持
gradient_checkpointing: false # 冻结模型不需要
compile: false # PyTorch 2.x compile,可加速但首次编译慢
# Data config
data:
image_size: 448 # 当 use_processor=false 时生效(Qwen3-VL 默认 448)
fps: 5 # BridgeData 采样率
num_workers: 1 # DataLoader workers(VideoCache 每个 worker 独立,太多爆 RAM)
pin_memory: true
shuffle: true
action_normalize: true # 对 action 做均值方差标准化
# 采样策略
frame_sampling: "all" # "all"=每帧都训练, "uniform"=均匀采样N帧
max_frames_per_episode: null # 仅 frame_sampling=uniform 时生效
max_cache_episodes: 150 # VideoCache 上限(RAM 有限时降低此值)
# Logging
logging:
wandb_project: "lingarm-vla"
wandb_run_name: "qwen3vl-frozen-mlp-bridge2.6k"
log_interval: 50 # 每50步打印一次日志
# Device
device: "cuda"
seed: 42