Download scripts/test_ddp_vlm_variable_shape.py from Jack04810/agentic-rl-main: direct link, hf CLI and curl.
- Browser
- Download file 11.8 kB
-
https://huggingface.co/Jack04810/agentic-rl-main/resolve/main/scripts/test_ddp_vlm_variable_shape.py
- Command line
-
hf download hf://Jack04810/agentic-rl-main/scripts/test_ddp_vlm_variable_shape.py
-
curl -L -o test_ddp_vlm_variable_shape.py https://huggingface.co/Jack04810/agentic-rl-main/resolve/main/scripts/test_ddp_vlm_variable_shape.py
11.8 kB
| #!/usr/bin/env python3 | |
| """Stress native DDP with deliberately unequal real ChartQA image lengths. | |
| The prior formal OPD run failed in DeepSpeed ZeRO-2's gradient-partition | |
| collective after one rank processed a 17-patch ChartQA image while the other | |
| ranks had much shorter inputs. This bounded regression test gives rank 0 the | |
| two largest hard-correct images and every other rank two of the smallest ones. | |
| It executes the same eight-microbatch gradient-accumulation pattern as the | |
| formal recipe, using DDP ``no_sync`` for the first seven microbatches. | |
| It intentionally exercises only the student VLM forward/backward. The full | |
| OPD smoke covers local frozen-teacher scoring; separating them makes a | |
| collective failure attributable to the distributed gradient path. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import contextlib | |
| import math | |
| import os | |
| import sys | |
| import time | |
| from datetime import timedelta | |
| from pathlib import Path | |
| import torch | |
| import torch.distributed as dist | |
| from PIL import Image | |
| from torch.nn.parallel import DistributedDataParallel | |
| from transformers import AutoProcessor, LlavaOnevisionForConditionalGeneration | |
| from transformers.models.llava_onevision.modeling_llava_onevision import image_size_to_num_patches | |
| PROJECT_ROOT = Path(__file__).resolve().parents[1] | |
| if str(PROJECT_ROOT) not in sys.path: | |
| sys.path.insert(0, str(PROJECT_ROOT)) | |
| from config.loader import load_config | |
| from data_utils.chart.data_collector import prepare_chart_rl_data | |
| from data_utils.chart.teacher_gate import filter_chartqa_rows_with_teacher_gate | |
| from data_utils.commom_util import collate_fn | |
| from opsd_utils.deepspeed_utils import student_forward_chunk_size | |
| from opsd_utils.teacher_batching import student_batch_num_images_tensor | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--config", default="config/config_opd_only_7b_chartqa.yaml") | |
| parser.add_argument("--updates", type=int, default=2) | |
| parser.add_argument("--gradient-accumulation-steps", type=int, default=8) | |
| parser.add_argument("--samples-per-rank", type=int, default=2) | |
| parser.add_argument( | |
| "--student-forward-chunk-size", | |
| type=int, | |
| default=0, | |
| help=( | |
| "Override the local student-forward chunk size. Zero uses the production " | |
| "student_forward_chunk_size() policy." | |
| ), | |
| ) | |
| parser.add_argument("--sync-each-chunk", action="store_true") | |
| parser.add_argument("--trace-chunks", action="store_true") | |
| parser.add_argument("--process-group-timeout-seconds", type=int, default=120) | |
| return parser.parse_args() | |
| def _memory_snapshot(device: torch.device) -> str: | |
| free_bytes, total_bytes = torch.cuda.mem_get_info(device) | |
| gib = 1024**3 | |
| return ( | |
| f"allocated_gib={torch.cuda.memory_allocated(device) / gib:.3f} " | |
| f"reserved_gib={torch.cuda.memory_reserved(device) / gib:.3f} " | |
| f"max_allocated_gib={torch.cuda.max_memory_allocated(device) / gib:.3f} " | |
| f"free_gib={free_bytes / gib:.3f} total_gib={total_bytes / gib:.3f}" | |
| ) | |
| def _rank_rows_with_extreme_shapes( | |
| config: dict, | |
| *, | |
| rank: int, | |
| world_size: int, | |
| samples_per_rank: int, | |
| ) -> tuple[list[dict], list[int]]: | |
| """Give rank 0 the largest images and all other ranks the smallest ones.""" | |
| rows = prepare_chart_rl_data(config["dataset"]["train_dataset"]) | |
| rows, _ = filter_chartqa_rows_with_teacher_gate(rows, config["dataset"]["teacher_gate"]) | |
| processor = AutoProcessor.from_pretrained( | |
| config["model"]["pretrained_model_path"], local_files_only=True | |
| ) | |
| grid = processor.image_processor.image_grid_pinpoints | |
| patch_size = processor.image_processor.size["height"] | |
| ranked: list[tuple[int, int, dict]] = [] | |
| for index, row in enumerate(rows): | |
| with Image.open(row["image"]) as image: | |
| patches = image_size_to_num_patches(image.size, grid, patch_size) | |
| ranked.append((int(patches), index, row)) | |
| ranked.sort(key=lambda item: (item[0], item[1])) | |
| if rank == 0: | |
| chosen = list(reversed(ranked[-samples_per_rank:])) | |
| else: | |
| start = (rank - 1) * samples_per_rank | |
| chosen = ranked[start : start + samples_per_rank] | |
| if len(chosen) != samples_per_rank: | |
| raise RuntimeError( | |
| "not enough ChartQA rows for distributed variable-shape DDP diagnostic" | |
| ) | |
| return [row for _patches, _index, row in chosen], [patches for patches, _index, _row in chosen] | |
| def _assert_extreme_shape_assignment( | |
| local_patch_counts: list[int], device: torch.device, rank: int, world_size: int) -> None: | |
| """Fail closed if the test no longer creates the intended rank skew.""" | |
| local_max = torch.tensor([max(local_patch_counts)], device=device, dtype=torch.long) | |
| gathered = [torch.zeros_like(local_max) for _ in range(world_size)] | |
| dist.all_gather(gathered, local_max) | |
| maxima = [int(item.item()) for item in gathered] | |
| if rank == 0: | |
| print(f"DDP variable-shape assignment: per_rank_max_patches={maxima}", flush=True) | |
| if maxima[0] <= max(maxima[1:]): | |
| raise RuntimeError( | |
| "rank 0 did not receive a strictly larger image than all other ranks: " | |
| f"{maxima}" | |
| ) | |
| def main() -> None: | |
| args = parse_args() | |
| if args.updates <= 0 or args.gradient_accumulation_steps <= 0 or args.samples_per_rank <= 0: | |
| raise ValueError("--updates, --gradient-accumulation-steps, and --samples-per-rank must be positive") | |
| local_rank = int(os.environ["LOCAL_RANK"]) | |
| torch.cuda.set_device(local_rank) | |
| if args.process_group_timeout_seconds <= 0: | |
| raise ValueError("--process-group-timeout-seconds must be positive") | |
| dist.init_process_group( | |
| backend="nccl", | |
| timeout=timedelta(seconds=args.process_group_timeout_seconds), | |
| ) | |
| rank = dist.get_rank() | |
| world_size = dist.get_world_size() | |
| device = torch.device("cuda", local_rank) | |
| try: | |
| config = load_config(args.config) | |
| local_rows, local_patch_counts = _rank_rows_with_extreme_shapes( | |
| config, | |
| rank=rank, | |
| world_size=world_size, | |
| samples_per_rank=args.samples_per_rank, | |
| ) | |
| _assert_extreme_shape_assignment(local_patch_counts, device, rank, world_size) | |
| model_path = config["model"]["pretrained_model_path"] | |
| processor = AutoProcessor.from_pretrained(model_path, local_files_only=True) | |
| processor.tokenizer.padding_side = "left" | |
| model = LlavaOnevisionForConditionalGeneration.from_pretrained( | |
| model_path, | |
| torch_dtype=torch.bfloat16, | |
| attn_implementation="flash_attention_2", | |
| local_files_only=True, | |
| low_cpu_mem_usage=True, | |
| ) | |
| model.base_model.vision_tower.requires_grad_(False) | |
| model.config.use_cache = False | |
| model.to(device) | |
| ddp_model = DistributedDataParallel( | |
| model, | |
| device_ids=[local_rank], | |
| output_device=local_rank, | |
| find_unused_parameters=True, | |
| ) | |
| optimizer = torch.optim.AdamW( | |
| (parameter for parameter in ddp_model.parameters() if parameter.requires_grad), | |
| lr=1.0e-6, | |
| ) | |
| batch = collate_fn(local_rows, processor, label_id=None) | |
| input_ids = batch["input_ids"].to(device) | |
| attention_mask = batch["attention_mask"].to(device) | |
| pixel_values = batch["pixel_values"].to(device=device, dtype=torch.bfloat16) | |
| image_sizes = batch["image_sizes"].to(device) | |
| batch_num_images = student_batch_num_images_tensor(pixel_values, input_ids.size(0)) | |
| production_chunk_size = student_forward_chunk_size( | |
| int(input_ids.size(0)), has_vision=True | |
| ) | |
| chunk_size = int(args.student_forward_chunk_size or production_chunk_size) | |
| if chunk_size <= 0 or chunk_size > int(input_ids.size(0)): | |
| raise ValueError( | |
| "student forward chunk size must be within the local batch: " | |
| f"chunk_size={chunk_size}, batch_size={input_ids.size(0)}" | |
| ) | |
| if rank == 0: | |
| print( | |
| "DDP variable-shape forward policy: " | |
| f"production_chunk_size={production_chunk_size} active_chunk_size={chunk_size} " | |
| f"forwards_per_backward={math.ceil(input_ids.size(0) / chunk_size)}", | |
| flush=True, | |
| ) | |
| dist.barrier() | |
| started = time.perf_counter() | |
| for update in range(args.updates): | |
| optimizer.zero_grad(set_to_none=True) | |
| for micro_step in range(args.gradient_accumulation_steps): | |
| is_sync_step = micro_step == args.gradient_accumulation_steps - 1 | |
| sync_context = contextlib.nullcontext() if is_sync_step else ddp_model.no_sync() | |
| with sync_context: | |
| chunk_losses = [] | |
| for chunk_start in range(0, int(input_ids.size(0)), chunk_size): | |
| chunk_end = min(chunk_start + chunk_size, int(input_ids.size(0))) | |
| if args.trace_chunks: | |
| print( | |
| "DDP chunk enter: " | |
| f"rank={rank}/{world_size} update={update} micro_step={micro_step} " | |
| f"chunk={chunk_start}:{chunk_end} {_memory_snapshot(device)}", | |
| flush=True, | |
| ) | |
| outputs = ddp_model( | |
| input_ids=input_ids[chunk_start:chunk_end], | |
| attention_mask=attention_mask[chunk_start:chunk_end], | |
| pixel_values=pixel_values[chunk_start:chunk_end], | |
| image_sizes=image_sizes[chunk_start:chunk_end], | |
| batch_num_images=batch_num_images[chunk_start:chunk_end], | |
| logits_to_keep=2, | |
| ) | |
| if args.sync_each_chunk: | |
| torch.cuda.synchronize(device) | |
| row_fraction = (chunk_end - chunk_start) / int(input_ids.size(0)) | |
| chunk_losses.append(outputs.logits.float().mean() * row_fraction) | |
| if args.trace_chunks: | |
| print( | |
| "DDP chunk done: " | |
| f"rank={rank}/{world_size} update={update} micro_step={micro_step} " | |
| f"chunk={chunk_start}:{chunk_end} {_memory_snapshot(device)}", | |
| flush=True, | |
| ) | |
| loss = sum(chunk_losses) / args.gradient_accumulation_steps | |
| loss.backward() | |
| optimizer.step() | |
| torch.cuda.synchronize(device) | |
| if rank == 0: | |
| print( | |
| "DDP variable-shape update complete: " | |
| f"update={update + 1}/{args.updates} loss={float(loss.detach().item()):.6f}", | |
| flush=True, | |
| ) | |
| dist.barrier() | |
| torch.cuda.synchronize(device) | |
| elapsed = time.perf_counter() - started | |
| print( | |
| "DDP VLM variable-shape stress passed: " | |
| f"rank={rank}/{world_size} patch_counts={local_patch_counts} " | |
| f"input_shape={tuple(input_ids.shape)} updates={args.updates} " | |
| f"ga={args.gradient_accumulation_steps} elapsed_s={elapsed:.3f}", | |
| flush=True, | |
| ) | |
| dist.barrier() | |
| finally: | |
| if dist.is_initialized(): | |
| dist.destroy_process_group() | |
| if __name__ == "__main__": | |
| main() | |