| 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 typing import Any, Dict, List |
| 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 flax.training.common_utils import get_metrics, onehot, shard |
| from transformers.models.vit.modeling_flax_vit import ViTConfig, FlaxViTForImageClassification |
| from model import OldLMCFlaxViTForImageClassification |
| from lmc_model import LMCFlaxViTForImageClassification, print_model, print_model_with_prefix |
| from datasets import build_dataset |
| import multiprocessing as mp |
| from pprint import pprint |
| 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 prepare_image_batch(images:torch.Tensor,labels:torch.Tensor) -> Dict[str, Any]: |
| images, labels = jnp.array(images),jnp.array(labels) |
| return {'images': shard(images),'labels': shard(labels)} |
| def accuracy(logits, labels, topk=(1,)): |
| maxk = max(topk) |
| batch_size = labels.shape[0] |
| topk_preds = jnp.argsort(logits, axis=-1)[:, -maxk:][:, ::-1] |
| res = [] |
| for k in topk: |
| correct = (topk_preds[:, :k] == labels[:, None]) |
| correct = jnp.any(correct, axis=1) |
| correct = jnp.sum(correct) |
| res.append(100.0 * correct / batch_size) |
| return res |
| def pretrained_new2old(new_params,old_params,config): |
| new_params = unfreeze(new_params) |
| old_params = unfreeze(old_params) |
| print_model(new_params) |
| print_model(old_params) |
| |
| new_params["vit"]["embeddings"] = copy.deepcopy(old_params["vit"]["embeddings"]) |
| new_params["vit"]["layernorm"] = copy.deepcopy(old_params["vit"]["layernorm"]) |
| new_params["classifier"] = copy.deepcopy(old_params["classifier"]) |
| |
| for i in range(config.num_hidden_layers): |
| str_i = str(i) |
| ref_layer = old_params["vit"]["encoder"]["layer"][str_i] |
| target_layer = new_params["vit"]["encoder"]["layer"][str_i] |
| |
| target_layer["layernorm_before"] = copy.deepcopy(ref_layer["layernorm_before"]) |
| target_layer["layernorm_after"] = copy.deepcopy(ref_layer["layernorm_after"]) |
| target_layer["moe"]['shared_experts']['intermediate'] = copy.deepcopy(ref_layer["mlp"]['intermediate']['dense']) |
| target_layer["moe"]['shared_experts']['output'] = copy.deepcopy(ref_layer["mlp"]['output']['dense']) |
| |
| target_layer["attention"] = copy.deepcopy(ref_layer["attention"]) |
| return freeze(new_params) |
| def main(args: argparse.Namespace): |
| train_loader, val_loader = imagenet_data_loader(args) |
| save_path = "/mnt/data/vinhbk/weights/imagenet/lr0.0005-rope-epochs300-batch256-shared1-routed0-topk0/" |
| old_config = ViTConfig.from_pretrained('/mnt/data/vinhbk/weights/imagenet/lr0.0005-rope-epochs300-batch256/config.json') |
| old_model = OldLMCFlaxViTForImageClassification(old_config,dtype=jnp.bfloat16) |
| new_config = copy.deepcopy(old_config) |
| new_config.position_embeddings = old_config.position_embeddings |
| new_config.rotary_value = old_config.rotary_value |
| new_config.q_lora_rank = 8 |
| new_config.qk_rope_head_dim = 64 |
| new_config.kv_lora_rank = 8 |
| new_config.v_head_dim = 64 |
| new_config.qk_nope_head_dim = 64 |
| new_config.attention_bias = True |
| new_config.routed_scaling_factor = 1.0 |
| new_config.lmc_layer_indices = [] |
| new_config.num_shared_experts = 1 |
| new_config.num_routed_experts = 0 |
| new_config.topk = 0 |
| |
| lr_schedule = optax.warmup_cosine_decay_schedule( |
| init_value=1e-6, |
| peak_value=5e-4, |
| warmup_steps=5*5004, |
| decay_steps=300*5004, |
| end_value=1e-5, |
| ) |
| |
| tx = optax.adamw( |
| learning_rate=lr_schedule, |
| b1=0.9, |
| b2=0.999, |
| eps=1e-8, |
| weight_decay=0.01, |
| ) |
| old_state = train_state.TrainState.create(apply_fn=old_model.__call__, params=old_model.params, tx=tx) |
| old_params = old_state.params |
| new_model = LMCFlaxViTForImageClassification(new_config,input_shape=(1,new_config.image_size, new_config.image_size, new_config.num_channels),seed=args.seed,dtype=jnp.bfloat16) |
| new_model.params = pretrained_new2old(new_params=copy.deepcopy(new_model.params),old_params=copy.deepcopy(old_state.params),config=new_model.config) |
| |
| |
| new_state = train_state.TrainState.create(apply_fn=new_model.__call__,params=new_model.params,tx=tx,) |
| print(old_state.params['classifier']['bias']) |
| print(new_state.params['classifier']['bias']) |
| |
| os.makedirs(save_path,exist_ok=True) |
| new_model.config.save_pretrained(save_path) |
| |
| def train_step(state, batch, rng): |
| dropout_rng, new_dropout_rng = jax.random.split(rng) |
| def loss_fn(params): |
| outputs = state.apply_fn(params=params,pixel_values=batch["images"],train=True,dropout_rng=dropout_rng,) |
| logits = outputs.logits if hasattr(outputs, "logits") else outputs[0] |
| loss = optax.softmax_cross_entropy_with_integer_labels(logits, batch["labels"]).mean() |
| return loss, logits |
| (loss, logits), grads = jax.value_and_grad(loss_fn, has_aux=True)(state.params) |
| grads = jax.lax.pmean(grads, axis_name="batch") |
| state = state.apply_gradients(grads=grads) |
| acc1, acc5 = accuracy(logits, batch["labels"], topk=(1, 5)) |
| metrics = {"loss": loss,"acc1": acc1,"acc5": acc5,"learning_rate": lr_schedule(state.step),} |
| metrics = jax.lax.pmean(metrics, axis_name="batch") |
| return state, metrics, new_dropout_rng |
| def eval_step(state, batch): |
| outputs = state.apply_fn(params=state.params,pixel_values=batch["images"],train=False,) |
| logits = outputs.logits if hasattr(outputs, "logits") else outputs[0] |
| loss = optax.softmax_cross_entropy_with_integer_labels(logits, batch["labels"]).mean() |
| acc1, acc5 = accuracy(logits, batch["labels"], topk=(1, 5)) |
| metrics = {"loss": loss,"acc1": acc1,"acc5": acc5,} |
| metrics = jax.lax.pmean(metrics, axis_name="batch") |
| return metrics |
| parallel_train_step = jax.pmap(train_step, "batch") |
| parallel_eval_step = jax.pmap(eval_step, "batch") |
| state = replicate(new_state) |
| rng = jax.random.PRNGKey(0) |
| train_metrics_stack = [] |
| train_loss = 0.0 |
| best_val_acc1 = 0.0 |
| |
| eval_results = [] |
| pbar = tqdm(enumerate(val_loader), desc="Evaluating", leave=False) |
| for batch_idx, (images, labels) in pbar: |
| batch = prepare_image_batch(images, labels) |
| |
| eval_metric = parallel_eval_step(state, batch) |
| eval_results.append(eval_metric) |
| |
| eval_metrics = get_metrics(eval_results) |
| eval_metrics = unreplicate(eval_metrics) |
| eval_metrics = jax.tree_util.tree_map(lambda x: x.mean(), eval_metrics) |
| val_loss, val_acc1, val_acc5 = float(eval_metrics["loss"]), float(eval_metrics["acc1"]), float(eval_metrics["acc5"]) |
| print("-" * 100) |
| print(f"valid loss {val_loss:5.4f} | valid acc@1 {val_acc1:6.2f}% | valid acc@5 {val_acc5:6.2f}%") |
| print("-" * 100) |
| if __name__ =="__main__": |
| parser = argparse.ArgumentParser(description="Fine-tune ViT with MoE on Imagenet") |
| parser.add_argument("--batch-size", type=int, default=256) |
| parser.add_argument("--data-path", type=str, required=True) |
| parser.add_argument('--data-set', default='IMNET', choices=['CIFAR', 'IMNET', 'INAT', 'INAT19']) |
| parser.add_argument("--input-size", type=int, default=224) |
| parser.add_argument('--num_workers', type=int, default=8) |
| parser.add_argument('--pin-mem', action='store_true') |
| parser.add_argument('--seed', type=int, default=0) |
| parser.add_argument('--color-jitter', type=float, default=0.4) |
| parser.add_argument('--aa', type=str, default='rand-m9-mstd0.5-inc1') |
| parser.add_argument('--train-interpolation', type=str, default='bicubic') |
| parser.add_argument('--reprob', type=float, default=0.25) |
| parser.add_argument('--remode', type=str, default='pixel') |
| parser.add_argument('--recount', type=int, default=1) |
| parser.add_argument("--dtype", choices=["float32", "float16", "bfloat16"], default="bfloat16", help="model datatype") |
| args = parser.parse_args() |
| main(args) |