import argparse import math import os import time from pathlib import Path import torch import torch.distributed as dist import torch.nn as nn import torch.nn.functional as F from torch.utils.data import DataLoader from torchvision import transforms from datasets import load_dataset from torch.optim import AdamW from tqdm import tqdm import json from datasets import DownloadConfig IMAGENET_MEAN = (0.485, 0.456, 0.406) IMAGENET_STD = (0.229, 0.224, 0.225) NUM_CLASSES = 1000 hf_token = "your hf token" # --------------------------------------------------------------------------- # # data # --------------------------------------------------------------------------- # def build_transforms(train, size=224): if train: return transforms.Compose([ transforms.RandomResizedCrop(size, scale=(0.08, 1.0), interpolation=transforms.InterpolationMode.BILINEAR), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD), ]) return transforms.Compose([ transforms.Resize(int(size * 256 / 224), interpolation=transforms.InterpolationMode.BILINEAR), transforms.CenterCrop(size), transforms.ToTensor(), transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD), ]) def make_collate(tf): """HF gives PIL images. Some ImageNet JPEGs are grayscale/CMYK -> force RGB.""" def collate(batch): pixels = torch.stack([tf(ex["image"].convert("RGB")) for ex in batch]) labels = torch.tensor([ex["label"] for ex in batch], dtype=torch.long) return pixels, labels return collate def build_loaders(image_size,batch_size,workers): # Transform train_tf = build_transforms(True, image_size) val_tf = build_transforms(False, image_size) # Streaming Dataset ds_train = load_dataset( "ILSVRC/imagenet-1k", split="train", streaming=True, token=hf_token, download_config=DownloadConfig(max_retries=20) ) ds_val = load_dataset( "ILSVRC/imagenet-1k", split="validation", streaming=True, token=hf_token, download_config=DownloadConfig(max_retries=20) ) # DataLoader train_loader = DataLoader( ds_train, batch_size=batch_size, num_workers=workers, pin_memory=True, collate_fn=make_collate(train_tf), prefetch_factor=5 ) val_loader = DataLoader( ds_val, batch_size=batch_size, num_workers=workers, pin_memory=True, collate_fn=make_collate(val_tf), prefetch_factor=5 ) return train_loader, val_loader class vitModel(nn.Module): def __init__(self,d_model,n_layers): super().__init__() self.positional_embedding = nn.Parameter(torch.zeros(784,d_model,requires_grad=True)) self.features = nn.Sequential( nn.Conv2d(3, 16, 3, stride=2, padding=1), # 224 -> 112 nn.ReLU(), nn.Conv2d(16, 32, 3, stride=2, padding=1), # 112 -> 56 nn.ReLU(), nn.Conv2d(32, 64, 3, stride=2, padding=1), # 56 -> 28 nn.ReLU(), nn.Conv2d(64, d_model, 3, padding=1), # 28 -> 28 ) self.encoder_layer = nn.TransformerEncoderLayer(d_model=d_model, nhead=d_model//16, batch_first=True) self.transformer_encoder = nn.TransformerEncoder(self.encoder_layer, num_layers=n_layers) self.fc = nn.Linear(d_model,1000) def forward(self, x): x = self.features(x) x = x.flatten(2) # (B, d_model, 1024) x = x.transpose(1, 2) # (B, 1024, d_model) x = x +self.positional_embedding x = self.transformer_encoder(x) x = torch.relu(x) x = x[:, -1, :] x = self.fc(x) return x def train(): # ---------------------------- # Hyperparameter # ---------------------------- image_size = 224 batch_size = 128 # "micro batch size" grad_accum_steps = 4 # real batch size = batch_size * grad_accum_steps workers = 8 d_model = 256 n_layers = 6 epochs = 200 lr = 3e-4 resume_checkpoint = None # example: "logs/123456_epoch_005_checkpoint.pth" device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # ---------------------------- # HF Hub push config # ---------------------------- push_to_hub = True hf_repo_id = "your repo" hf_token = "your hf token" if push_to_hub: from huggingface_hub import HfApi, create_repo create_repo(hf_repo_id, exist_ok=True, token=hf_token) hf_api = HfApi(token=hf_token) # ---------------------------- # Model # ---------------------------- model = vitModel( d_model=d_model, n_layers=n_layers, ).to(device) num_params = sum(p.numel() for p in model.parameters()) # ---------------------------- # Logging # ---------------------------- log_dir = "logs" os.makedirs(log_dir, exist_ok=True) config_path = os.path.join(log_dir, f"{num_params}_config.json") loss_log_path = os.path.join(log_dir, f"{num_params}_loss.jsonl") config = { "image_size": image_size, "batch_size": batch_size, "grad_accum_steps": grad_accum_steps, "effective_batch_size": batch_size * grad_accum_steps, "workers": workers, "d_model": d_model, "n_layers": n_layers, "epochs": epochs, "lr": lr, "resume_checkpoint": resume_checkpoint, "device": str(device), "num_parameters": num_params, } if not os.path.exists(config_path): with open(config_path, "w") as f: json.dump(config, f, indent=4) if not os.path.exists(loss_log_path): open(loss_log_path, "w").close() # ---------------------------- # Data # ---------------------------- train_loader, val_loader = build_loaders( image_size=image_size, batch_size=batch_size, workers=workers, ) # ---------------------------- # Optimizer # ---------------------------- optimizer = AdamW(model.parameters(), lr=lr) criterion = nn.CrossEntropyLoss() start_epoch = 0 global_step = 0 # ---------------------------- # Resume # ---------------------------- if resume_checkpoint is not None: checkpoint = torch.load(resume_checkpoint, map_location=device) model.load_state_dict(checkpoint["model"]) optimizer.load_state_dict(checkpoint["optimizer"]) start_epoch = checkpoint["epoch"] global_step = checkpoint["global_step"] # ---------------------------- # Train # ---------------------------- for epoch in range(start_epoch, epochs): model.train() running_loss = 0 running_correct = 0 total = 0 optimizer.zero_grad() pbar = tqdm(train_loader) for step, (images, labels) in enumerate(pbar): images = images.to(device, non_blocking=True) labels = labels.to(device, non_blocking=True) outputs = model(images) loss = criterion(outputs, labels) # gradient accumulation (loss / grad_accum_steps).backward() # ---------------------------- # optimizer step # grad_accum step # ---------------------------- if (step + 1) % grad_accum_steps == 0: optimizer.step() optimizer.zero_grad() # ---------------------------- # JSONL Logging # ---------------------------- with open(loss_log_path, "a") as f: json.dump({ "step": global_step, "epoch": epoch + 1, "loss": float(loss.item()) }, f) f.write("\n") global_step += 1 running_loss += loss.item() pred = outputs.argmax(dim=1) running_correct += ( pred == labels ).sum().item() total += labels.size(0) pbar.set_description( f"Epoch {epoch+1}/{epochs}" ) pbar.set_postfix( loss=f"{running_loss/(step+1):.4f}", acc=f"{100*running_correct/total:.2f}%" ) # --------------------------------- # Processing remaining gradients at the end of the epoch # --------------------------------- if (step + 1) % grad_accum_steps != 0: optimizer.step() optimizer.zero_grad() # ---------------------------- # Validation # ---------------------------- model.eval() val_correct = 0 val_total = 0 with torch.no_grad(): for images, labels in val_loader: images = images.to(device, non_blocking=True) labels = labels.to(device, non_blocking=True) outputs = model(images) pred = outputs.argmax(dim=1) val_correct += ( pred == labels ).sum().item() val_total += labels.size(0) train_acc = 100 * running_correct / total val_acc = 100 * val_correct / val_total print( f"Epoch {epoch+1} " f"Train Acc: {train_acc:.2f}% " f"Val Acc: {val_acc:.2f}%" ) # ---------------------------- # Checkpoint save # ---------------------------- checkpoint_filename = ( f"epoch_{epoch+1:03d}_valacc_{val_acc:.2f}.pth" ) checkpoint_path = os.path.join( log_dir, checkpoint_filename ) torch.save({ "epoch": epoch + 1, "global_step": global_step, "model": model.state_dict(), "optimizer": optimizer.state_dict(), "train_acc": train_acc, "val_acc": val_acc, }, checkpoint_path) # ---------------------------- # HF Hub push # ---------------------------- if push_to_hub: try: hf_api.upload_file( path_or_fileobj=checkpoint_path, path_in_repo=checkpoint_filename, repo_id=hf_repo_id, token=hf_token, ) hf_api.upload_file( path_or_fileobj=config_path, path_in_repo=os.path.basename(config_path), repo_id=hf_repo_id, token=hf_token, ) hf_api.upload_file( path_or_fileobj=loss_log_path, path_in_repo=os.path.basename(loss_log_path), repo_id=hf_repo_id, token=hf_token, ) print( f"[HF Hub] pushed {checkpoint_filename} -> {hf_repo_id}" ) except Exception as e: print( f"[HF Hub] push failed at epoch {epoch+1}: {e}" ) if __name__ == "__main__": train()