import os, time, torch, nvtx import torch.nn as nn import torch.optim as optim import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data.distributed import DistributedSampler from torch.utils.data import Dataset, DataLoader from torchvision.models import vit_l_32 WORLD_SIZE = int(os.environ.get("WORLD_SIZE", 1)) BATCH_SIZE = 32 IMG_SIZE = 224 WARMUP_STEPS = 10 PROFILE_STEPS = 3 COOLDOWN_STEPS = 1 TOTAL_STEPS = WARMUP_STEPS + PROFILE_STEPS + COOLDOWN_STEPS N_WORKERS = 8 def get_model(): return vit_l_32(weights=None) # A synthetic dataset with random images and labels class FakeDataset(Dataset): def __len__(self): return TOTAL_STEPS * BATCH_SIZE * WORLD_SIZE def __getitem__(self, index): img = torch.randn((3, IMG_SIZE, IMG_SIZE)) label = torch.randint(0, 1000, (1,)).item() return img, label def get_data_iter(rank, world_size): dataset = FakeDataset() sampler = DistributedSampler(dataset, num_replicas=world_size, rank=rank, shuffle=True) train_loader = DataLoader(dataset, batch_size=BATCH_SIZE, sampler=sampler, num_workers=N_WORKERS, pin_memory=True) return iter(train_loader)