Files
VLA/train_qact.py
T

119 lines
3.0 KiB
Python
Raw Normal View History

2025-09-07 14:59:17 +08:00
import os
import json
import time
import yaml
import wandb
from argparse import ArgumentParser
2025-09-11 13:18:33 +08:00
from accelerate import (
Accelerator,
DistributedDataParallelKwargs,
DataLoaderConfiguration,
)
2025-09-07 14:59:17 +08:00
from wall_x.trainer.qwen_vl_act_trainer import QwenVlAct_Trainer
def setup_environment():
"""Set up environment variables for training."""
os.environ["TOKENIZERS_PARALLELISM"] = "false"
def load_config(config_path):
"""Load configuration from YAML file."""
with open(config_path, "r") as f:
config = yaml.load(f, Loader=yaml.FullLoader)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Set model_type in data config if not already set
config["data"]["model_type"] = config.get("model_type")
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
return config
def setup_accelerator(config):
"""Initialize and configure the accelerator for distributed training."""
2025-09-11 13:18:33 +08:00
print(
f"[{time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())}] Preparing accelerator"
)
2025-09-07 14:59:17 +08:00
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
accelerator_dataloader_config = DataLoaderConfiguration(dispatch_batches=False)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
accelerator = Accelerator(
kwargs_handlers=[ddp_kwargs],
mixed_precision="bf16",
dataloader_config=accelerator_dataloader_config,
2025-09-11 13:18:33 +08:00
gradient_accumulation_steps=config.get("gradient_accumulation_steps", 1),
2025-09-07 14:59:17 +08:00
)
2025-09-11 13:18:33 +08:00
print(
f"[{time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())}] Accelerator initialization complete"
)
2025-09-07 14:59:17 +08:00
return accelerator
def setup_logging(config, accelerator):
"""Set up logging with wandb for the main process."""
if not accelerator.is_main_process:
return None
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Create save directory if it doesn't exist
save_path = config["save_path"]
if not os.path.exists(save_path):
print(f"Save path {save_path} does not exist, creating directory.")
os.makedirs(save_path, exist_ok=True)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
print("Configuration:")
print("=" * 50)
print(json.dumps(config, indent=2, ensure_ascii=False))
print("=" * 50)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Initialize wandb logger
logger = wandb.init(
project=config["log_project"],
name=config["log_name"],
save_code=False,
force=False,
)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
return logger
def main(args):
"""Main training function."""
setup_environment()
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Load configuration
config = load_config(args.config)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Set up accelerator
accelerator = setup_accelerator(config)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Set up logging
logger = setup_logging(config, accelerator)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Initialize trainer
trainer = QwenVlAct_Trainer(
config=config,
logger=logger,
accelerator=accelerator,
seed=args.seed,
data_config_path=args.config,
)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Start training
trainer.fit()
2025-09-11 13:18:33 +08:00
if __name__ == "__main__":
2025-09-07 14:59:17 +08:00
parser = ArgumentParser(description="Training script for Wall-X model")
2025-09-11 13:18:33 +08:00
parser.add_argument(
"--config", type=str, required=True, help="Path to configuration YAML file"
)
parser.add_argument(
"--seed", type=int, default=42, help="Random seed for reproducibility"
)
2025-09-07 14:59:17 +08:00
args = parser.parse_args()
2025-09-11 13:18:33 +08:00
main(args)