MCPcopy Create free account
hub / github.com/SHAILAB-IPEC/OpenFly-Platform / main

Function main

train/train.py:89–199  ·  view source on GitHub ↗
(data_args=None, training_args=None)

Source from the content-addressed store, hash-verified

87
88
89def 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,

Callers 1

train.pyFile · 0.70

Calls 15

run_setupMethod · 0.95
run_vla_trainingMethod · 0.95
finalizeMethod · 0.95
initialize_overwatchFunction · 0.90
load_vlaFunction · 0.90
save_dataset_statisticsFunction · 0.90
TrainingStrategyClass · 0.90
VLAMetricsClass · 0.90
sumFunction · 0.85
local_rankMethod · 0.80
get_image_transformMethod · 0.80

Tested by

no test coverage detected