(data_args=None, training_args=None)
| 87 | |
| 88 | |
| 89 | def main(data_args=None, training_args=None): |
| 90 | |
| 91 | # Initialize Overwatch =>> Wraps `logging.Logger` |
| 92 | overwatch = initialize_overwatch(__name__) |
| 93 | |
| 94 | torch.cuda.set_device(device_id := overwatch.local_rank()) |
| 95 | torch.cuda.empty_cache() |
| 96 | run_id = ( |
| 97 | f"n{training_args.expected_world_size // 8}+b{training_args.per_device_batch_size}+x{training_args.seed}" |
| 98 | ) |
| 99 | |
| 100 | |
| 101 | worker_init_fn = None |
| 102 | |
| 103 | os.makedirs(run_dir := (data_args.run_root_dir / run_id), exist_ok=True) |
| 104 | os.makedirs(data_args.run_root_dir / run_id / "checkpoints", exist_ok=True) |
| 105 | |
| 106 | model = load_vla(training_args.pretrained_checkpoint, hf_token=training_args.hf_token, load_for_training=True, grid_size=training_args.grid_size) |
| 107 | |
| 108 | for param in model.parameters(): |
| 109 | assert param.dtype == torch.float32, f"Loaded VLM parameter not in full precision: {param}" |
| 110 | |
| 111 | # Determine training "stage" based on frozen vs unfrozen parameters --> supports different fine-tuning schemes! |
| 112 | if not training_args.freeze_vision_backbone and not training_args.freeze_llm_backbone: |
| 113 | stage = "vla-full-train" # Full fine-tuning |
| 114 | elif training_args.freeze_vision_backbone and not training_args.freeze_llm_backbone: |
| 115 | stage = "vla-train" # Frozen vision encoder |
| 116 | elif not training_args.freeze_vision_backbone and training_args.freeze_llm_backbone: |
| 117 | assert training_args.unfreeze_last_llm_layer, "You should unfreeze at least the last layer of your LLM!" |
| 118 | stage = "vla-sandwich-train" # Fine-tuning vision encoder, projector, and LLM last layer |
| 119 | elif training_args.freeze_vision_backbone and training_args.freeze_llm_backbone: |
| 120 | assert training_args.unfreeze_last_llm_layer, "Need to unfreeze at least last LLM layer to train!" |
| 121 | stage = "vla-last-layer-train" # Fine-tuning LLM last layer only |
| 122 | else: |
| 123 | raise ValueError( |
| 124 | "Weight freezing configuration not supported. VLA config has the following parameters: " |
| 125 | f"freeze_vision_backbone: {training_args.freeze_vision_backbone}" |
| 126 | f"freeze_llm_backbone: {training_args.freeze_llm_backbone}" |
| 127 | f"unfreeze_last_llm_layer: {training_args.unfreeze_last_llm_layer}" |
| 128 | ) |
| 129 | |
| 130 | # [Explicit] Call to `freeze_backbones` here for clarity =>> will log exactly what is/is not frozen |
| 131 | overwatch.info(f"Stage Info: ") |
| 132 | model.freeze_backbones(stage) |
| 133 | |
| 134 | # Print number of total/trainable model parameters |
| 135 | num_params = sum(p.numel() for p in model.parameters()) |
| 136 | num_trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) |
| 137 | overwatch.info( |
| 138 | f"# Parameters (in millions): {num_params / 10**6:.3f} Total, {num_trainable_params / 10**6:.3f} Trainable" |
| 139 | ) |
| 140 | |
| 141 | vla_dataset, action_tokenizer, collator = get_vla_dataset_and_collator( |
| 142 | data_args.data_root_dir, |
| 143 | data_args.data_mix, |
| 144 | image_transform=model.vision_backbone.get_image_transform(), |
| 145 | tokenizer=model.llm_backbone.get_tokenizer(), |
| 146 | default_image_resolution=model.vision_backbone.default_image_resolution, |
no test coverage detected