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)