| import argparse |
| import os |
| import jax |
| import flax |
| import copy |
| import json |
| import optax |
| import torch |
| import wandb |
| import numpy as np |
| from tqdm import tqdm |
| import jax.numpy as jnp |
| from flax import linen as nn |
| from flax.jax_utils import replicate, unreplicate |
| from flax.training import checkpoints, train_state |
| from flax.core.frozen_dict import freeze, unfreeze |
| from flax.traverse_util import flatten_dict, unflatten_dict |
| from transformers.models.vit.modeling_flax_vit import ViTConfig, FlaxViTForImageClassification |
| from flax.training.common_utils import get_metrics, onehot, shard |
| from lmc_model import LMCFlaxViTForImageClassification, print_model, print_model_with_prefix |
| from datasets import build_dataset |
| import multiprocessing as mp |
| from jax import debug |
| from pprint import pprint |
| from typing import Any, Dict, List |
| import shutil |
| mp.set_start_method("spawn", force=True) |
| os.environ["WANDB_API_KEY"] = "fc72050bcc0dc7f7502b5416938f8bd0c4b30fc7" |
| def imagenet_data_loader(args): |
| dataset_train, args.nb_classes = build_dataset(is_train=True, args=args) |
| dataset_val, _ = build_dataset(is_train=False, args=args) |
| sampler_train = torch.utils.data.RandomSampler(dataset_train) |
| sampler_val = torch.utils.data.SequentialSampler(dataset_val) |
| data_loader_train = torch.utils.data.DataLoader( |
| dataset_train, sampler=sampler_train, batch_size=args.batch_size, |
| num_workers=args.num_workers, pin_memory=args.pin_mem, drop_last=True |
| ) |
| data_loader_val = torch.utils.data.DataLoader( |
| dataset_val, sampler=sampler_val, batch_size=args.batch_size, |
| num_workers=args.num_workers, pin_memory=args.pin_mem, drop_last=False |
| ) |
| return data_loader_train, data_loader_val |
| def main(args: argparse.Namespace): |
| save_path = "/mnt/data/vinhbk/weights/imagenet/lr0.0005-rope-epochs300-batch256-shared1-routed0-topk0" |
| model = LMCFlaxViTForImageClassification.from_pretrained("/mnt/data/vinhbk/weights/imagenet/lr0.0005-rope-epochs300-batch256-shared1-routed0-topk0/temp_1291032") |
| step = 1291032 |
| num_global_steps = 5004*args.epochs |
| num_warmup_steps = 5004*args.warmup_epochs |
| lr_schedule =optax.warmup_cosine_decay_schedule( |
| init_value=args.warmup_lr, |
| peak_value=args.lr, |
| warmup_steps=num_warmup_steps, |
| decay_steps=num_global_steps, |
| end_value=args.min_lr, |
| ) |
| tx = optax.adamw( |
| learning_rate=lr_schedule, |
| b1=args.adamw_beta1, |
| b2=args.adamw_beta2, |
| eps=args.adamw_eps, |
| weight_decay=args.weight_decay, |
| ) |
| state = train_state.TrainState.create(apply_fn=model.__call__, params=model.params, tx=tx) |
| state = state.replace(step=step) |
| checkpoints.save_checkpoint(ckpt_dir=save_path,target=state,step=state.step,prefix="last_",keep=1) |
|
|
|
|
| if __name__ == "__main__": |
| parser = argparse.ArgumentParser(description="Fine-tune ViT with MoE on Imagenet") |
| parser.add_argument("--epochs", type=int, default=30) |
| parser.add_argument("--batch-size", type=int, default=64) |
| parser.add_argument("--lr", type=float, default=5e-4) |
| parser.add_argument("--weight-decay", type=float, default=0.01) |
| parser.add_argument('--sched', default='cosine', type=str, metavar='SCHEDULER') |
| parser.add_argument('--warmup-lr', type=float, default=1e-6, metavar='LR', help='warmup learning rate (default: 1e-6)') |
| parser.add_argument('--warmup-epochs', type=int, default=3, metavar='N',help='epochs to warmup LR, if scheduler supports') |
| parser.add_argument('--min-lr', type=float, default=1e-5, metavar='LR',help='lower lr bound for cyclic schedulers that hit 0 (1e-5)') |
| parser.add_argument("--adamw-beta1", type=float, default=0.9) |
| parser.add_argument("--adamw-beta2", type=float, default=0.999) |
| parser.add_argument("--adamw-eps", type=float, default=1e-8) |
| parser.add_argument("--patience", type=int, default=10, help="Early stopping patience") |
| args = parser.parse_args() |
| main(args) |