From 7b34818951e2b8ab3d9e0ea7e3550468142db1aa Mon Sep 17 00:00:00 2001 From: init-22 Date: Wed, 6 Aug 2025 16:19:02 +0000 Subject: [PATCH 1/7] spacing added for testing --- README.md | 1 + 1 file changed, 1 insertion(+) diff --git a/README.md b/README.md index 0666d21d5..525dc9210 100644 --- a/README.md +++ b/README.md @@ -37,6 +37,7 @@ Submissions are evaluated based on their "time-to-result", i.e., the wall-clock --- > [!IMPORTANT] + > For future iterations of the AlgoPerf: Training Algorithms benchmark competition, we are switching to a rolling leaderboard, making a few changes to the competition rules, and also run all selected submissions on our hardware. **To submit your algorithm to the next iteration of the benchmark, please see our [How to Submit](#how-to-submit) section and the [submission repository](https://github.com/mlcommons/submissions_algorithms) which hosts the up to date AlgoPerf leaderboard.** ## Table of Contents From 5741c76898a80e51637758c154b9656c9660bd79 Mon Sep 17 00:00:00 2001 From: init-22 Date: Wed, 6 Aug 2025 16:55:55 +0000 Subject: [PATCH 2/7] copying schedule free code --- .../criteo1tb/criteo1tb_jax/workload.py | 2 + .../criteo1tb/criteo1tb_pytorch/workload.py | 1 + custom_pytorch_jax_converter.py | 99 ++++++ .../schedule_free/jax/submission.py | 227 +++++++++++++ .../schedule_free/pytorch/submission.py | 319 ++++++++++++++++++ 5 files changed, 648 insertions(+) create mode 100644 custom_pytorch_jax_converter.py create mode 100644 reference_algorithms/paper_baselines/schedule_free/jax/submission.py create mode 100644 reference_algorithms/paper_baselines/schedule_free/pytorch/submission.py diff --git a/algoperf/workloads/criteo1tb/criteo1tb_jax/workload.py b/algoperf/workloads/criteo1tb/criteo1tb_jax/workload.py index 283b3be8e..e82c5bdf1 100644 --- a/algoperf/workloads/criteo1tb/criteo1tb_jax/workload.py +++ b/algoperf/workloads/criteo1tb/criteo1tb_jax/workload.py @@ -11,6 +11,7 @@ from algoperf import param_utils, spec from algoperf.workloads.criteo1tb.criteo1tb_jax import models from algoperf.workloads.criteo1tb.workload import BaseCriteo1TbDlrmSmallWorkload +from custom_pytorch_jax_converter import use_pytorch_weights class Criteo1TbDlrmSmallWorkload(BaseCriteo1TbDlrmSmallWorkload): @@ -104,6 +105,7 @@ def init_model_fn( jnp.ones(input_shape, jnp.float32), ) initial_params = initial_variables['params'] + initial_params = use_pytorch_weights(file_name="~/results/pytorch_base_model_criteo1tb_1_july.pth") self._param_shapes = param_utils.jax_param_shapes(initial_params) self._param_types = param_utils.jax_param_types(self._param_shapes) return jax_utils.replicate(initial_params), None diff --git a/algoperf/workloads/criteo1tb/criteo1tb_pytorch/workload.py b/algoperf/workloads/criteo1tb/criteo1tb_pytorch/workload.py index 74f91de43..ba17e86cf 100644 --- a/algoperf/workloads/criteo1tb/criteo1tb_pytorch/workload.py +++ b/algoperf/workloads/criteo1tb/criteo1tb_pytorch/workload.py @@ -82,6 +82,7 @@ def init_model_fn(self, rng: spec.RandomState) -> spec.ModelInitState: use_layer_norm=self.use_layer_norm, embedding_init_multiplier=self.embedding_init_multiplier, ) + torch.save(model.state_dict(), '~/results/pytorch_base_model_criteo1tb_1_july.pth') self._param_shapes = param_utils.pytorch_param_shapes(model) self._param_types = param_utils.pytorch_param_types(self._param_shapes) model.to(DEVICE) diff --git a/custom_pytorch_jax_converter.py b/custom_pytorch_jax_converter.py new file mode 100644 index 000000000..e5c723919 --- /dev/null +++ b/custom_pytorch_jax_converter.py @@ -0,0 +1,99 @@ +import torch +import numpy as np +import jax +import jax.numpy as jnp +import logging +import copy +import copy +from jax.tree_util import tree_map + + +def use_pytorch_weights(file_name: str): + """ + Jax default parameter structure: + dict_keys(['Dense_0', 'Dense_1', 'Dense_2', 'Dense_3', 'Dense_4', 'Dense_5', 'Dense_6', 'Dense_7', 'embedding_table']) + + Pytorch stateduct structure: + dict_keys(['embedding_chunk_0', 'embedding_chunk_1', 'embedding_chunk_2', 'embedding_chunk_3', 'bot_mlp.0.weight', 'bot_mlp.0.bias', 'bot_mlp.2.weight', 'bot_mlp.2.bias', 'bot_mlp.4.weight', 'bot_mlp.4.bias', 'top_mlp.0.weight', 'top_mlp.0.bias', 'top_mlp.2.weight', 'top_mlp.2.bias', 'top_mlp.4.weight', 'top_mlp.4.bias', 'top_mlp.6.weight', 'top_mlp.6.bias', 'top_mlp.8.weight', 'top_mlp.8.bias']) + + + The following function converts the PyTorch weights to the Jax format + """ + + jax_copy = {} + + # Load PyTorch state_dict lazily to CPU + state_dict = torch.load(file_name, map_location='cpu') + print(state_dict.keys()) + + # Convert PyTorch tensors to NumPy arrays + numpy_weights = {k: v.cpu().numpy() for k, v in state_dict.items()} + + # --- Embedding Table --- + embedding_table = np.concatenate([ + numpy_weights[f'embedding_chunk_{i}'] for i in range(4) + ], axis=0) # adjust axis if chunking is not vertical + + jax_copy['embedding_table'] = jnp.array(embedding_table) + + # --- Bot MLP: Dense_0 to Dense_2 --- + for i, j in zip([0, 2, 4], range(3)): + jax_copy[f'Dense_{j}'] = {} + jax_copy[f'Dense_{j}']['kernel'] = jnp.array(numpy_weights[f'bot_mlp.{i}.weight'].T) + jax_copy[f'Dense_{j}']['bias'] = jnp.array(numpy_weights[f'bot_mlp.{i}.bias']) + + # --- Top MLP: Dense_3 to Dense_7 --- + for i, j in zip([0, 2, 4, 6, 8], range(3, 8)): + jax_copy[f'Dense_{j}'] = {} + jax_copy[f'Dense_{j}']['kernel'] = jnp.array(numpy_weights[f'top_mlp.{i}.weight'].T) + jax_copy[f'Dense_{j}']['bias'] = jnp.array(numpy_weights[f'top_mlp.{i}.bias']) + + del state_dict + return jax_copy + + +def maybe_unreplicate(pytree): + """If leading axis matches device count, strip it assuming it's pmap replication.""" + num_devices = jax.device_count() + return jax.tree_util.tree_map( + lambda x: x[0] if isinstance(x, jax.Array) and x.shape[0] == num_devices else x, + pytree + ) + + +def move_to_cpu(tree): + return jax.tree_util.tree_map(lambda x: jax.device_put(x, device=jax.devices("cpu")[0]), tree) + + +def are_weights_equal(params1, params2, atol=1e-6, rtol=1e-6): + """Compares two JAX PyTrees of weights and logs where they differ, safely handling PMAP replication.""" + # Attempt to unreplicate if needed + + params1 = maybe_unreplicate(params1) + params2 = maybe_unreplicate(params2) + + params1 = move_to_cpu(params1) + params2 = move_to_cpu(params2) + + all_equal = True + + def compare_fn(p1, p2): + nonlocal all_equal + if not jnp.allclose(p1, p2, atol=atol, rtol=rtol): + logging.info("❌ Mismatch found:") + logging.info(f"Shape : {p1.shape}, Shape 2: {p2.shape}") + logging.info(f"Max diff: {jnp.max(jnp.abs(p1 - p2))}") + all_equal = False + return jnp.allclose(p1, p2, atol=atol, rtol=rtol) + + try: + jax.tree_util.tree_map(compare_fn, params1, params2) + except Exception as e: + logging.info("❌ Structure mismatch or error during comparison:", exc_info=True) + return False + + if all_equal: + logging.info("✅ All weights are equal (within tolerance)") + del params1 + del params2 + return all_equal \ No newline at end of file diff --git a/reference_algorithms/paper_baselines/schedule_free/jax/submission.py b/reference_algorithms/paper_baselines/schedule_free/jax/submission.py new file mode 100644 index 000000000..95cd71339 --- /dev/null +++ b/reference_algorithms/paper_baselines/schedule_free/jax/submission.py @@ -0,0 +1,227 @@ +"""Submission file for an Schedule Free AdamW optimizer in Jax.""" + +import functools +from typing import Dict, Iterator, List, Tuple +import optax + +from flax import jax_utils +import jax +from jax import lax +import jax.numpy as jnp +from optax.contrib import schedule_free_adamw +from algoperf import spec +from custom_pytorch_jax_converter import use_pytorch_weights, are_weights_equal + +_GRAD_CLIP_EPS = 1e-6 + +HPARAMS = { + "dropout_rate": 0.1, + "learning_rate": 0.0025, + "one_minus_beta1": 0.1, + "beta2": 0.9955159689799007, + "weight_decay": 0.08121616522670176, + "warmup_factor": 0.02, + "weight_lr_power": 2, + "label_smoothing": 0.2, + "r": 0.75, + "eps": 1e-8, +} + +def init_optimizer_state(workload: spec.Workload, + model_params: spec.ParameterContainer, + model_state: spec.ModelAuxiliaryState, + hyperparameters: spec.Hyperparameters, + rng: spec.RandomState) -> spec.OptimizerState: + """Creates an AdamW optimizer and a learning rate schedule.""" + model_params + del model_state + del rng + + opt_init_fn, opt_update_fn = schedule_free_adamw( + learning_rate=HPARAMS['learning_rate'], + warmup_steps=int(HPARAMS['warmup_factor'] * workload.step_hint * 0.75), + + b1=1.0 - HPARAMS['one_minus_beta1'], + b2=HPARAMS['beta2'], + eps=HPARAMS['eps'], + weight_decay=HPARAMS['weight_decay'], + weight_lr_power=HPARAMS['weight_lr_power'], + # state_dtype=jnp.bfloat16 + ) + + model_params = jax_utils.unreplicate(model_params) + optimizer_state = opt_init_fn(model_params) + + return jax_utils.replicate(optimizer_state), opt_update_fn + + +@functools.partial( + jax.pmap, + axis_name='batch', + in_axes=(None, None, 0, 0, 0, 0, 0, None, None), + static_broadcasted_argnums=(0, 1), + donate_argnums=(2, 3, 4)) +def pmapped_train_step(workload, + opt_update_fn, + model_state, + optimizer_state, + current_param_container, + batch, + rng, + grad_clip, + label_smoothing): + + def _loss_fn(params): + """Loss function used for training.""" + logits, new_model_state = workload.model_fn( + params, + batch, + model_state, + spec.ForwardPassMode.TRAIN, + rng, + update_batch_norm=True) + loss_dict = workload.loss_fn( + label_batch=batch['targets'], + logits_batch=logits, + mask_batch=batch.get('weights'), + label_smoothing=label_smoothing) + summed_loss = loss_dict['summed'] + n_valid_examples = loss_dict['n_valid_examples'] + return summed_loss, (n_valid_examples, new_model_state) + + grad_fn = jax.value_and_grad(_loss_fn, has_aux=True) + (summed_loss, (n_valid_examples, new_model_state)), grad = grad_fn( + current_param_container) + # Get correct global mean loss and grad. + (summed_loss, n_valid_examples, grad) = lax.psum( + (summed_loss, n_valid_examples, grad), axis_name='batch') + loss = summed_loss / n_valid_examples + grad = jax.tree_map(lambda x: x / n_valid_examples, grad) + + grad_norm = jnp.sqrt( + sum(jnp.sum(g**2) for g in jax.tree_util.tree_leaves(grad))) + + # Extract the leaves of the pytree + leaves = jax.tree_util.tree_leaves(grad) + # Count the total number of elements in all leaves + total_size = sum(jnp.size(leaf) for leaf in leaves) + + # jax.debug.print('GRAD NORM {}', grad_norm) + # jax.debug.print('NUM PARAMS {}', total_size) + + if grad_clip is not None: + grad_scaling_factor = grad_clip / (grad_norm + _GRAD_CLIP_EPS) + grad_scaling_factor = jax.lax.clamp(min=0.0, x=grad_scaling_factor, max=1.0) + grad = jax.tree_map(lambda x: x * grad_scaling_factor, grad) + + updates, new_optimizer_state = opt_update_fn(grad, optimizer_state, + current_param_container) + updated_params = optax.apply_updates(current_param_container, updates) + return new_optimizer_state, updated_params, new_model_state, loss, grad_norm + + +def update_params(workload: spec.Workload, + current_param_container: spec.ParameterContainer, + current_params_types: spec.ParameterTypeTree, + model_state: spec.ModelAuxiliaryState, + hyperparameters: spec.Hyperparameters, + batch: Dict[str, spec.Tensor], + loss_type: spec.LossType, + optimizer_state: spec.OptimizerState, + eval_results: List[Tuple[int, float]], + global_step: int, + rng: spec.RandomState) -> spec.UpdateReturn: + """Return (updated_optimizer_state, updated_params, updated_model_state).""" + del current_params_types + del loss_type + del eval_results + + optimizer_state, opt_update_fn = optimizer_state + per_device_rngs = jax.random.split(rng, jax.local_device_count()) + if hasattr(hyperparameters, 'label_smoothing'): + label_smoothing = hyperparameters.label_smoothing + else: + label_smoothing = 0.0 + if hasattr(hyperparameters, 'grad_clip'): + grad_clip = hyperparameters.grad_clip + else: + grad_clip = None + outputs = pmapped_train_step(workload, + opt_update_fn, + model_state, + optimizer_state, + current_param_container, + batch, + per_device_rngs, + grad_clip, + label_smoothing) + new_optimizer_state, new_params, new_model_state, loss, grad_norm = outputs + + # Log loss, grad_norm. + if global_step % 100 == 0 and workload.metrics_logger is not None: + workload.metrics_logger.append_scalar_metrics( + { + 'loss': loss[0], + 'grad_norm': grad_norm[0], + }, global_step) + + # Log the number of parameters. + if global_step % 100 == 0: + date_ = "2025-07-01" + file_name = f"/results/schedule_free_pytorch_weights/criteo1tb_{date_}_after_{global_step}_steps.pth" + params = use_pytorch_weights(file_name=file_name) + are_weights_equal(new_params, params) + del params + + return (new_optimizer_state, opt_update_fn), new_params, new_model_state + + +def get_batch_size(workload_name): + # Return the global batch size. + if workload_name == 'criteo1tb': + return 262_144 + elif workload_name == 'fastmri': + return 32 + elif workload_name == 'imagenet_resnet': + return 1024 + elif workload_name == 'imagenet_resnet_silu': + return 512 + elif workload_name == 'imagenet_resnet_gelu': + return 512 + elif workload_name == 'imagenet_vit': + return 1024 + elif workload_name == 'librispeech_conformer': + return 256 + elif workload_name == 'librispeech_deepspeech': + return 256 + elif workload_name == 'ogbg': + return 512 + elif workload_name == 'wmt': + return 128 + elif workload_name == 'mnist': + return 16 + else: + raise ValueError(f'Unsupported workload name: {workload_name}.') + + +def data_selection(workload: spec.Workload, + input_queue: Iterator[Dict[str, spec.Tensor]], + optimizer_state: spec.OptimizerState, + current_param_container: spec.ParameterContainer, + model_state: spec.ModelAuxiliaryState, + hyperparameters: spec.Hyperparameters, + global_step: int, + rng: spec.RandomState) -> Dict[str, spec.Tensor]: + """Select data from the infinitely repeating, pre-shuffled input queue. + Each element of the queue is a batch of training examples and labels. + """ + del workload + del optimizer_state + del current_param_container + del model_state + del hyperparameters + del global_step + del rng + batch = next(input_queue) + return batch + diff --git a/reference_algorithms/paper_baselines/schedule_free/pytorch/submission.py b/reference_algorithms/paper_baselines/schedule_free/pytorch/submission.py new file mode 100644 index 000000000..450745737 --- /dev/null +++ b/reference_algorithms/paper_baselines/schedule_free/pytorch/submission.py @@ -0,0 +1,319 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +import math +from typing import Dict, Iterator, List, Tuple +from absl import logging +import torch +import torch.distributed.nn as dist_nn +from algoperf import spec +from algoperf.pytorch_utils import pytorch_setup + +USE_PYTORCH_DDP = pytorch_setup()[0] +HPARAMS = { + "dropout_rate": 0.1, + "learning_rate": 0.0025, + "one_minus_beta1": 0.1, + "beta2": 0.9955159689799007, + "weight_decay": 0.08121616522670176, + "warmup_factor": 0.02, + "weight_lr_power": 2, + "label_smoothing": 0.2, + "r": 0.75, + "conformer_bs": 192, +} + + +class AdamWScheduleFree(torch.optim.Optimizer): + r"""Schedule Free AdamW + """ + def __init__(self, params, + lr=1e-3, + betas=(0.9, 0.999), + eps=1e-8, + weight_decay=0, + weight_lr_power=2, + warmup_steps=0, + r=0, + ): + defaults = dict(lr=lr, + betas=betas, + eps=eps, + r=r, + k=0, + weight_sum=0.0, + lr_max=0.0, + warmup_steps=warmup_steps, + weight_lr_power=weight_lr_power, + weight_decay=weight_decay) + + super().__init__(params, defaults) + + def reset(self): + for group in self.param_groups: + group['k'] = 0 + group['lr_max'] = 0 + group['weight_sum'] = 0 + + for p in group['params']: + # State initialization + state = self.state[p] + state['z'].copy_(state['x0']) + p.data.copy_(state['x0']) + state['exp_avg_sq'].zero_() + + def step(self, closure): + """Performs a single optimization step. + + Arguments: + closure (callable, optional): A closure that reevaluates the model + and returns the loss. + """ + # Swap to extrapolated point: + for group in self.param_groups: + beta1, beta2 = group['betas'] + r = group['r'] + k = group['k'] + + for p in group['params']: + # State initialization + state = self.state[p] + if 'z' not in state: + state['z'] = torch.clone(p.data) + state['exp_avg_sq'] = torch.zeros_like(p.data, dtype=torch.bfloat16) + state['x0'] = p.data.cpu() + + z = state['z'] + + # Extrapolate + #p = p + (1-beta1)*(z-p) + #p.data.mul_(beta1).add_(z, alpha=1-beta1) + p.data.lerp_(end=z, weight=1-beta1) + + # Evaluate gradient at extrapolated point + loss = closure() + + for group in self.param_groups: + eps = group['eps'] + k = group['k'] + warmup_steps = group['warmup_steps'] + + if k < warmup_steps: + sched = (k+1) / warmup_steps + else: + sched = 1.0 + annealed_lr = group['lr']*sched + + lr = max(annealed_lr, eps) + + decay = group['weight_decay'] + beta1, beta2 = group['betas'] + weight_lr_power = group['weight_lr_power'] + + r = group['r'] + lr_max = group['lr_max'] = max(lr, group['lr_max']) + + weight = ((k+1)**r) * (lr_max**weight_lr_power) + weight_sum = group['weight_sum'] = group['weight_sum'] + weight + + ckp1 = weight/weight_sum + + bias_correction2 = 1 - beta2 ** (k+1) + step_size = lr * math.sqrt(bias_correction2) + + for p in group['params']: + if p.grad is None: + continue + grad = p.grad.data + + state = self.state[p] + + exp_avg_sq = state['exp_avg_sq'] + z = state['z'] + + # Unextrapolate + #p = (p - (1-beta1)*z)/beta1 + #p.data.sub_(z, alpha=1-beta1).div_(beta1) + p.data.lerp_(end=z, weight=1-1/beta1) + + # Decay the first and second moment running average coefficient + exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1-beta2) + denom = exp_avg_sq.sqrt().add_(eps) + + z.addcdiv_(grad, denom, value=-step_size) + + # Decay + z.sub_(p.data, alpha=step_size*decay) + + ### Take step + #p.data.mul_(1-ckp1).add_(z, alpha=ckp1) + p.data.lerp_(end=z, weight=ckp1) + + group['k'] = k+1 + return loss + +def init_optimizer_state(workload: spec.Workload, + model_params: spec.ParameterContainer, + model_state: spec.ModelAuxiliaryState, + hyperparameters: spec.Hyperparameters, + rng: spec.RandomState) -> spec.OptimizerState: + del model_state + + optimizer = AdamWScheduleFree( + model_params.parameters(), + lr=HPARAMS['learning_rate'], + betas=(1.0 - HPARAMS['one_minus_beta1'], HPARAMS['beta2']), + warmup_steps=int(HPARAMS['warmup_factor'] * workload.step_hint * 0.75), + weight_decay=HPARAMS['weight_decay'], + weight_lr_power=HPARAMS['weight_lr_power'], + r=HPARAMS['r']) + + optimizer_state = {'optimizer':optimizer, 'max_checked_eval_step': -1, 'has_forced_reset': False, 'first_eval': False, } + return optimizer_state + +def update_params(workload: spec.Workload, + current_param_container: spec.ParameterContainer, + current_params_types: spec.ParameterTypeTree, + model_state: spec.ModelAuxiliaryState, + hyperparameters: spec.Hyperparameters, + batch: Dict[str, spec.Tensor], + loss_type: spec.LossType, + optimizer_state: spec.OptimizerState, + eval_results: List[Tuple[int, float]], + global_step: int, + rng: spec.RandomState) -> spec.UpdateReturn: + """Return (updated_optimizer_state, updated_params, updated_model_state).""" + del current_params_types + del loss_type + del hyperparameters + + + metric_name = workload.target_metric_name + # TODO - remove force_reset + eval_step = len(eval_results) + if (global_step > workload.step_hint*0.10) and optimizer_state['max_checked_eval_step'] < eval_step: + optimizer_state['max_checked_eval_step'] = eval_step + # Don't do resetting on workloads that don't run eval often enough + if len(eval_results) >= 4: # and optimizer_state["first_eval_is_far_from_target"]: + val_metric = f"validation/{metric_name}" + initial_eval = eval_results[0][1][val_metric] + latest_eval = eval_results[-1][1][val_metric] + second_latest_eval = eval_results[-2][1][val_metric] + third_latest_eval = eval_results[-3][1][val_metric] + fourth_latest_eval = eval_results[-4][1][val_metric] + MARGIN = 0.01 + if metric_name in ["loss", "wer"]: + # Decreasing eval workloads should be flipped + initial_eval = -initial_eval + latest_eval = -latest_eval + second_latest_eval = -second_latest_eval + third_latest_eval = -third_latest_eval + fourth_latest_eval = -fourth_latest_eval + # Higher is better + # scale as a curve from 0 --> 1 + # if the eval values are far from the target (i.e. - worse than initial) and stays far from the target for 4 evals + if (latest_eval - initial_eval < MARGIN) and (latest_eval - second_latest_eval < MARGIN) and (second_latest_eval - third_latest_eval < MARGIN) and (third_latest_eval - fourth_latest_eval < MARGIN): + # Reset parameters since we appear to have diverged + logging.info("Reseting All Weights ") + logging.info(f"Global Step: {global_step}") + optimizer_state['has_forced_reset'] = True + + # Perform reset + del model_state + model_state = None + optimizer_state['optimizer'].reset() + + # Decrease learning rate by 2x if it diverged. + for param_group in optimizer_state['optimizer'].param_groups: + param_group['lr'] = param_group['lr']/2.0 + + ########### + + current_model = current_param_container + current_model.train() + + new_model_state = None + + def closure(): + nonlocal new_model_state + optimizer_state['optimizer'].zero_grad() + + logits_batch, new_model_state = workload.model_fn( + params=current_model, + augmented_and_preprocessed_input_batch=batch, + model_state=model_state, + mode=spec.ForwardPassMode.TRAIN, + rng=rng, + update_batch_norm=True) + + loss_dict = workload.loss_fn( + label_batch=batch['targets'], + logits_batch=logits_batch, + mask_batch=batch.get('weights'), + label_smoothing=HPARAMS['label_smoothing']) + summed_loss = loss_dict['summed'] + n_valid_examples = loss_dict['n_valid_examples'] + if USE_PYTORCH_DDP: + # Use dist_nn.all_reduce to ensure correct loss and gradient scaling. + summed_loss = dist_nn.all_reduce(summed_loss) + n_valid_examples = dist_nn.all_reduce(n_valid_examples) + loss = summed_loss / n_valid_examples + + loss.backward() + return loss + + loss = optimizer_state['optimizer'].step(closure) + + return (optimizer_state, current_param_container, new_model_state) + + +def get_batch_size(workload_name): + # Return the global batch size. + if workload_name == 'criteo1tb': + return 262_144 + elif workload_name == 'fastmri': + return 16 # 32 + elif workload_name == 'imagenet_resnet': + return 1024 + elif workload_name == 'imagenet_vit': + return 1024 + elif workload_name == 'librispeech_conformer': + return 224 + elif workload_name == 'librispeech_deepspeech': + return 128 # 256 + elif workload_name == 'ogbg': + return 512 + elif workload_name == 'wmt': + return 128 + elif workload_name == 'mnist': + return 16 + elif workload_name == 'imagenet_resnet_gelu': + return 512 + elif workload_name == 'imagenet_resnet_silu': + return 512 + else: + raise ValueError(f'Unsupported workload name: {workload_name}.') + +def data_selection(workload: spec.Workload, + input_queue: Iterator[Dict[str, spec.Tensor]], + optimizer_state: spec.OptimizerState, + current_param_container: spec.ParameterContainer, + model_state: spec.ModelAuxiliaryState, + hyperparameters: spec.Hyperparameters, + global_step: int, + rng: spec.RandomState) -> Dict[str, spec.Tensor]: + """Select data from the infinitely repeating, pre-shuffled input queue. + Each element of the queue is a batch of training examples and labels. + """ + del workload + del optimizer_state + del current_param_container + del model_state + del hyperparameters + del global_step + del rng + batch = next(input_queue) + return batch + From 73bcca46bb9c66b388e34629f19b09f800d76e8b Mon Sep 17 00:00:00 2001 From: init-22 Date: Fri, 8 Aug 2025 10:25:59 +0000 Subject: [PATCH 3/7] setting up to run the workloads --- .../criteo1tb/criteo1tb_jax/workload.py | 2 - .../criteo1tb/criteo1tb_pytorch/workload.py | 1 - custom_pytorch_jax_converter.py | 99 ------------------- 3 files changed, 102 deletions(-) delete mode 100644 custom_pytorch_jax_converter.py diff --git a/algoperf/workloads/criteo1tb/criteo1tb_jax/workload.py b/algoperf/workloads/criteo1tb/criteo1tb_jax/workload.py index e82c5bdf1..283b3be8e 100644 --- a/algoperf/workloads/criteo1tb/criteo1tb_jax/workload.py +++ b/algoperf/workloads/criteo1tb/criteo1tb_jax/workload.py @@ -11,7 +11,6 @@ from algoperf import param_utils, spec from algoperf.workloads.criteo1tb.criteo1tb_jax import models from algoperf.workloads.criteo1tb.workload import BaseCriteo1TbDlrmSmallWorkload -from custom_pytorch_jax_converter import use_pytorch_weights class Criteo1TbDlrmSmallWorkload(BaseCriteo1TbDlrmSmallWorkload): @@ -105,7 +104,6 @@ def init_model_fn( jnp.ones(input_shape, jnp.float32), ) initial_params = initial_variables['params'] - initial_params = use_pytorch_weights(file_name="~/results/pytorch_base_model_criteo1tb_1_july.pth") self._param_shapes = param_utils.jax_param_shapes(initial_params) self._param_types = param_utils.jax_param_types(self._param_shapes) return jax_utils.replicate(initial_params), None diff --git a/algoperf/workloads/criteo1tb/criteo1tb_pytorch/workload.py b/algoperf/workloads/criteo1tb/criteo1tb_pytorch/workload.py index ba17e86cf..74f91de43 100644 --- a/algoperf/workloads/criteo1tb/criteo1tb_pytorch/workload.py +++ b/algoperf/workloads/criteo1tb/criteo1tb_pytorch/workload.py @@ -82,7 +82,6 @@ def init_model_fn(self, rng: spec.RandomState) -> spec.ModelInitState: use_layer_norm=self.use_layer_norm, embedding_init_multiplier=self.embedding_init_multiplier, ) - torch.save(model.state_dict(), '~/results/pytorch_base_model_criteo1tb_1_july.pth') self._param_shapes = param_utils.pytorch_param_shapes(model) self._param_types = param_utils.pytorch_param_types(self._param_shapes) model.to(DEVICE) diff --git a/custom_pytorch_jax_converter.py b/custom_pytorch_jax_converter.py deleted file mode 100644 index e5c723919..000000000 --- a/custom_pytorch_jax_converter.py +++ /dev/null @@ -1,99 +0,0 @@ -import torch -import numpy as np -import jax -import jax.numpy as jnp -import logging -import copy -import copy -from jax.tree_util import tree_map - - -def use_pytorch_weights(file_name: str): - """ - Jax default parameter structure: - dict_keys(['Dense_0', 'Dense_1', 'Dense_2', 'Dense_3', 'Dense_4', 'Dense_5', 'Dense_6', 'Dense_7', 'embedding_table']) - - Pytorch stateduct structure: - dict_keys(['embedding_chunk_0', 'embedding_chunk_1', 'embedding_chunk_2', 'embedding_chunk_3', 'bot_mlp.0.weight', 'bot_mlp.0.bias', 'bot_mlp.2.weight', 'bot_mlp.2.bias', 'bot_mlp.4.weight', 'bot_mlp.4.bias', 'top_mlp.0.weight', 'top_mlp.0.bias', 'top_mlp.2.weight', 'top_mlp.2.bias', 'top_mlp.4.weight', 'top_mlp.4.bias', 'top_mlp.6.weight', 'top_mlp.6.bias', 'top_mlp.8.weight', 'top_mlp.8.bias']) - - - The following function converts the PyTorch weights to the Jax format - """ - - jax_copy = {} - - # Load PyTorch state_dict lazily to CPU - state_dict = torch.load(file_name, map_location='cpu') - print(state_dict.keys()) - - # Convert PyTorch tensors to NumPy arrays - numpy_weights = {k: v.cpu().numpy() for k, v in state_dict.items()} - - # --- Embedding Table --- - embedding_table = np.concatenate([ - numpy_weights[f'embedding_chunk_{i}'] for i in range(4) - ], axis=0) # adjust axis if chunking is not vertical - - jax_copy['embedding_table'] = jnp.array(embedding_table) - - # --- Bot MLP: Dense_0 to Dense_2 --- - for i, j in zip([0, 2, 4], range(3)): - jax_copy[f'Dense_{j}'] = {} - jax_copy[f'Dense_{j}']['kernel'] = jnp.array(numpy_weights[f'bot_mlp.{i}.weight'].T) - jax_copy[f'Dense_{j}']['bias'] = jnp.array(numpy_weights[f'bot_mlp.{i}.bias']) - - # --- Top MLP: Dense_3 to Dense_7 --- - for i, j in zip([0, 2, 4, 6, 8], range(3, 8)): - jax_copy[f'Dense_{j}'] = {} - jax_copy[f'Dense_{j}']['kernel'] = jnp.array(numpy_weights[f'top_mlp.{i}.weight'].T) - jax_copy[f'Dense_{j}']['bias'] = jnp.array(numpy_weights[f'top_mlp.{i}.bias']) - - del state_dict - return jax_copy - - -def maybe_unreplicate(pytree): - """If leading axis matches device count, strip it assuming it's pmap replication.""" - num_devices = jax.device_count() - return jax.tree_util.tree_map( - lambda x: x[0] if isinstance(x, jax.Array) and x.shape[0] == num_devices else x, - pytree - ) - - -def move_to_cpu(tree): - return jax.tree_util.tree_map(lambda x: jax.device_put(x, device=jax.devices("cpu")[0]), tree) - - -def are_weights_equal(params1, params2, atol=1e-6, rtol=1e-6): - """Compares two JAX PyTrees of weights and logs where they differ, safely handling PMAP replication.""" - # Attempt to unreplicate if needed - - params1 = maybe_unreplicate(params1) - params2 = maybe_unreplicate(params2) - - params1 = move_to_cpu(params1) - params2 = move_to_cpu(params2) - - all_equal = True - - def compare_fn(p1, p2): - nonlocal all_equal - if not jnp.allclose(p1, p2, atol=atol, rtol=rtol): - logging.info("❌ Mismatch found:") - logging.info(f"Shape : {p1.shape}, Shape 2: {p2.shape}") - logging.info(f"Max diff: {jnp.max(jnp.abs(p1 - p2))}") - all_equal = False - return jnp.allclose(p1, p2, atol=atol, rtol=rtol) - - try: - jax.tree_util.tree_map(compare_fn, params1, params2) - except Exception as e: - logging.info("❌ Structure mismatch or error during comparison:", exc_info=True) - return False - - if all_equal: - logging.info("✅ All weights are equal (within tolerance)") - del params1 - del params2 - return all_equal \ No newline at end of file From 41291e0e13d18cdf9967cb7e884b637a09da63c8 Mon Sep 17 00:00:00 2001 From: init-22 Date: Sun, 10 Aug 2025 13:36:51 +0000 Subject: [PATCH 4/7] removing unwated code which calls function for comparison --- .../paper_baselines/schedule_free/jax/submission.py | 8 -------- 1 file changed, 8 deletions(-) diff --git a/reference_algorithms/paper_baselines/schedule_free/jax/submission.py b/reference_algorithms/paper_baselines/schedule_free/jax/submission.py index 95cd71339..162dd76d9 100644 --- a/reference_algorithms/paper_baselines/schedule_free/jax/submission.py +++ b/reference_algorithms/paper_baselines/schedule_free/jax/submission.py @@ -10,7 +10,6 @@ import jax.numpy as jnp from optax.contrib import schedule_free_adamw from algoperf import spec -from custom_pytorch_jax_converter import use_pytorch_weights, are_weights_equal _GRAD_CLIP_EPS = 1e-6 @@ -165,13 +164,6 @@ def update_params(workload: spec.Workload, 'grad_norm': grad_norm[0], }, global_step) - # Log the number of parameters. - if global_step % 100 == 0: - date_ = "2025-07-01" - file_name = f"/results/schedule_free_pytorch_weights/criteo1tb_{date_}_after_{global_step}_steps.pth" - params = use_pytorch_weights(file_name=file_name) - are_weights_equal(new_params, params) - del params return (new_optimizer_state, opt_update_fn), new_params, new_model_state From 4f17b9675bb1b7446bf750637bf3710315e39a98 Mon Sep 17 00:00:00 2001 From: init-22 Date: Thu, 14 Aug 2025 17:29:27 +0000 Subject: [PATCH 5/7] fixing linting issues as per ruff --- .../paper_baselines/schedule_free/jax/submission.py | 12 +++++++----- .../schedule_free/pytorch/submission.py | 6 ++++-- 2 files changed, 11 insertions(+), 7 deletions(-) diff --git a/reference_algorithms/paper_baselines/schedule_free/jax/submission.py b/reference_algorithms/paper_baselines/schedule_free/jax/submission.py index 162dd76d9..f726bec79 100644 --- a/reference_algorithms/paper_baselines/schedule_free/jax/submission.py +++ b/reference_algorithms/paper_baselines/schedule_free/jax/submission.py @@ -2,13 +2,14 @@ import functools from typing import Dict, Iterator, List, Tuple -import optax -from flax import jax_utils import jax -from jax import lax import jax.numpy as jnp +import optax +from flax import jax_utils +from jax import lax from optax.contrib import schedule_free_adamw + from algoperf import spec _GRAD_CLIP_EPS = 1e-6 @@ -101,9 +102,10 @@ def _loss_fn(params): sum(jnp.sum(g**2) for g in jax.tree_util.tree_leaves(grad))) # Extract the leaves of the pytree - leaves = jax.tree_util.tree_leaves(grad) + # leaves = jax.tree_util.tree_leaves(grad) + # Count the total number of elements in all leaves - total_size = sum(jnp.size(leaf) for leaf in leaves) + # total_size = sum(jnp.size(leaf) for leaf in leaves) # jax.debug.print('GRAD NORM {}', grad_norm) # jax.debug.print('NUM PARAMS {}', total_size) diff --git a/reference_algorithms/paper_baselines/schedule_free/pytorch/submission.py b/reference_algorithms/paper_baselines/schedule_free/pytorch/submission.py index 450745737..d1d9c25b9 100644 --- a/reference_algorithms/paper_baselines/schedule_free/pytorch/submission.py +++ b/reference_algorithms/paper_baselines/schedule_free/pytorch/submission.py @@ -5,9 +5,11 @@ # LICENSE file in the root directory of this source tree. import math from typing import Dict, Iterator, List, Tuple -from absl import logging + import torch import torch.distributed.nn as dist_nn +from absl import logging + from algoperf import spec from algoperf.pytorch_utils import pytorch_setup @@ -264,7 +266,7 @@ def closure(): loss.backward() return loss - loss = optimizer_state['optimizer'].step(closure) + _ = optimizer_state['optimizer'].step(closure) return (optimizer_state, current_param_container, new_model_state) From abff710dbe4091d7680f0e220a6ae23c38a479b6 Mon Sep 17 00:00:00 2001 From: init-22 Date: Thu, 14 Aug 2025 17:43:07 +0000 Subject: [PATCH 6/7] fixing ruff formatting issues --- .../schedule_free/jax/submission.py | 192 +++++---- .../schedule_free/pytorch/submission.py | 400 ++++++++++-------- 2 files changed, 318 insertions(+), 274 deletions(-) diff --git a/reference_algorithms/paper_baselines/schedule_free/jax/submission.py b/reference_algorithms/paper_baselines/schedule_free/jax/submission.py index f726bec79..fa6f4062f 100644 --- a/reference_algorithms/paper_baselines/schedule_free/jax/submission.py +++ b/reference_algorithms/paper_baselines/schedule_free/jax/submission.py @@ -15,38 +15,40 @@ _GRAD_CLIP_EPS = 1e-6 HPARAMS = { - "dropout_rate": 0.1, - "learning_rate": 0.0025, - "one_minus_beta1": 0.1, - "beta2": 0.9955159689799007, - "weight_decay": 0.08121616522670176, - "warmup_factor": 0.02, - "weight_lr_power": 2, - "label_smoothing": 0.2, - "r": 0.75, - "eps": 1e-8, + 'dropout_rate': 0.1, + 'learning_rate': 0.0025, + 'one_minus_beta1': 0.1, + 'beta2': 0.9955159689799007, + 'weight_decay': 0.08121616522670176, + 'warmup_factor': 0.02, + 'weight_lr_power': 2, + 'label_smoothing': 0.2, + 'r': 0.75, + 'eps': 1e-8, } -def init_optimizer_state(workload: spec.Workload, - model_params: spec.ParameterContainer, - model_state: spec.ModelAuxiliaryState, - hyperparameters: spec.Hyperparameters, - rng: spec.RandomState) -> spec.OptimizerState: + +def init_optimizer_state( + workload: spec.Workload, + model_params: spec.ParameterContainer, + model_state: spec.ModelAuxiliaryState, + hyperparameters: spec.Hyperparameters, + rng: spec.RandomState, +) -> spec.OptimizerState: """Creates an AdamW optimizer and a learning rate schedule.""" model_params del model_state del rng opt_init_fn, opt_update_fn = schedule_free_adamw( - learning_rate=HPARAMS['learning_rate'], - warmup_steps=int(HPARAMS['warmup_factor'] * workload.step_hint * 0.75), - - b1=1.0 - HPARAMS['one_minus_beta1'], - b2=HPARAMS['beta2'], - eps=HPARAMS['eps'], - weight_decay=HPARAMS['weight_decay'], - weight_lr_power=HPARAMS['weight_lr_power'], - # state_dtype=jnp.bfloat16 + learning_rate=HPARAMS['learning_rate'], + warmup_steps=int(HPARAMS['warmup_factor'] * workload.step_hint * 0.75), + b1=1.0 - HPARAMS['one_minus_beta1'], + b2=HPARAMS['beta2'], + eps=HPARAMS['eps'], + weight_decay=HPARAMS['weight_decay'], + weight_lr_power=HPARAMS['weight_lr_power'], + # state_dtype=jnp.bfloat16 ) model_params = jax_utils.unreplicate(model_params) @@ -56,50 +58,57 @@ def init_optimizer_state(workload: spec.Workload, @functools.partial( - jax.pmap, - axis_name='batch', - in_axes=(None, None, 0, 0, 0, 0, 0, None, None), - static_broadcasted_argnums=(0, 1), - donate_argnums=(2, 3, 4)) -def pmapped_train_step(workload, - opt_update_fn, - model_state, - optimizer_state, - current_param_container, - batch, - rng, - grad_clip, - label_smoothing): - + jax.pmap, + axis_name='batch', + in_axes=(None, None, 0, 0, 0, 0, 0, None, None), + static_broadcasted_argnums=(0, 1), + donate_argnums=(2, 3, 4), +) +def pmapped_train_step( + workload, + opt_update_fn, + model_state, + optimizer_state, + current_param_container, + batch, + rng, + grad_clip, + label_smoothing, +): def _loss_fn(params): """Loss function used for training.""" logits, new_model_state = workload.model_fn( - params, - batch, - model_state, - spec.ForwardPassMode.TRAIN, - rng, - update_batch_norm=True) + params, + batch, + model_state, + spec.ForwardPassMode.TRAIN, + rng, + update_batch_norm=True, + ) loss_dict = workload.loss_fn( - label_batch=batch['targets'], - logits_batch=logits, - mask_batch=batch.get('weights'), - label_smoothing=label_smoothing) + label_batch=batch['targets'], + logits_batch=logits, + mask_batch=batch.get('weights'), + label_smoothing=label_smoothing, + ) summed_loss = loss_dict['summed'] n_valid_examples = loss_dict['n_valid_examples'] return summed_loss, (n_valid_examples, new_model_state) grad_fn = jax.value_and_grad(_loss_fn, has_aux=True) (summed_loss, (n_valid_examples, new_model_state)), grad = grad_fn( - current_param_container) + current_param_container + ) # Get correct global mean loss and grad. (summed_loss, n_valid_examples, grad) = lax.psum( - (summed_loss, n_valid_examples, grad), axis_name='batch') + (summed_loss, n_valid_examples, grad), axis_name='batch' + ) loss = summed_loss / n_valid_examples grad = jax.tree_map(lambda x: x / n_valid_examples, grad) grad_norm = jnp.sqrt( - sum(jnp.sum(g**2) for g in jax.tree_util.tree_leaves(grad))) + sum(jnp.sum(g**2) for g in jax.tree_util.tree_leaves(grad)) + ) # Extract the leaves of the pytree # leaves = jax.tree_util.tree_leaves(grad) @@ -115,23 +124,26 @@ def _loss_fn(params): grad_scaling_factor = jax.lax.clamp(min=0.0, x=grad_scaling_factor, max=1.0) grad = jax.tree_map(lambda x: x * grad_scaling_factor, grad) - updates, new_optimizer_state = opt_update_fn(grad, optimizer_state, - current_param_container) + updates, new_optimizer_state = opt_update_fn( + grad, optimizer_state, current_param_container + ) updated_params = optax.apply_updates(current_param_container, updates) return new_optimizer_state, updated_params, new_model_state, loss, grad_norm -def update_params(workload: spec.Workload, - current_param_container: spec.ParameterContainer, - current_params_types: spec.ParameterTypeTree, - model_state: spec.ModelAuxiliaryState, - hyperparameters: spec.Hyperparameters, - batch: Dict[str, spec.Tensor], - loss_type: spec.LossType, - optimizer_state: spec.OptimizerState, - eval_results: List[Tuple[int, float]], - global_step: int, - rng: spec.RandomState) -> spec.UpdateReturn: +def update_params( + workload: spec.Workload, + current_param_container: spec.ParameterContainer, + current_params_types: spec.ParameterTypeTree, + model_state: spec.ModelAuxiliaryState, + hyperparameters: spec.Hyperparameters, + batch: Dict[str, spec.Tensor], + loss_type: spec.LossType, + optimizer_state: spec.OptimizerState, + eval_results: List[Tuple[int, float]], + global_step: int, + rng: spec.RandomState, +) -> spec.UpdateReturn: """Return (updated_optimizer_state, updated_params, updated_model_state).""" del current_params_types del loss_type @@ -147,25 +159,28 @@ def update_params(workload: spec.Workload, grad_clip = hyperparameters.grad_clip else: grad_clip = None - outputs = pmapped_train_step(workload, - opt_update_fn, - model_state, - optimizer_state, - current_param_container, - batch, - per_device_rngs, - grad_clip, - label_smoothing) + outputs = pmapped_train_step( + workload, + opt_update_fn, + model_state, + optimizer_state, + current_param_container, + batch, + per_device_rngs, + grad_clip, + label_smoothing, + ) new_optimizer_state, new_params, new_model_state, loss, grad_norm = outputs # Log loss, grad_norm. if global_step % 100 == 0 and workload.metrics_logger is not None: workload.metrics_logger.append_scalar_metrics( - { - 'loss': loss[0], - 'grad_norm': grad_norm[0], - }, global_step) - + { + 'loss': loss[0], + 'grad_norm': grad_norm[0], + }, + global_step, + ) return (new_optimizer_state, opt_update_fn), new_params, new_model_state @@ -198,14 +213,16 @@ def get_batch_size(workload_name): raise ValueError(f'Unsupported workload name: {workload_name}.') -def data_selection(workload: spec.Workload, - input_queue: Iterator[Dict[str, spec.Tensor]], - optimizer_state: spec.OptimizerState, - current_param_container: spec.ParameterContainer, - model_state: spec.ModelAuxiliaryState, - hyperparameters: spec.Hyperparameters, - global_step: int, - rng: spec.RandomState) -> Dict[str, spec.Tensor]: +def data_selection( + workload: spec.Workload, + input_queue: Iterator[Dict[str, spec.Tensor]], + optimizer_state: spec.OptimizerState, + current_param_container: spec.ParameterContainer, + model_state: spec.ModelAuxiliaryState, + hyperparameters: spec.Hyperparameters, + global_step: int, + rng: spec.RandomState, +) -> Dict[str, spec.Tensor]: """Select data from the infinitely repeating, pre-shuffled input queue. Each element of the queue is a batch of training examples and labels. """ @@ -218,4 +235,3 @@ def data_selection(workload: spec.Workload, del rng batch = next(input_queue) return batch - diff --git a/reference_algorithms/paper_baselines/schedule_free/pytorch/submission.py b/reference_algorithms/paper_baselines/schedule_free/pytorch/submission.py index d1d9c25b9..dae84863c 100644 --- a/reference_algorithms/paper_baselines/schedule_free/pytorch/submission.py +++ b/reference_algorithms/paper_baselines/schedule_free/pytorch/submission.py @@ -1,6 +1,6 @@ # Copyright (c) Meta Platforms, Inc. and affiliates. # All rights reserved. -# +# # This source code is licensed under the license found in the # LICENSE file in the root directory of this source tree. import math @@ -15,152 +15,159 @@ USE_PYTORCH_DDP = pytorch_setup()[0] HPARAMS = { - "dropout_rate": 0.1, - "learning_rate": 0.0025, - "one_minus_beta1": 0.1, - "beta2": 0.9955159689799007, - "weight_decay": 0.08121616522670176, - "warmup_factor": 0.02, - "weight_lr_power": 2, - "label_smoothing": 0.2, - "r": 0.75, - "conformer_bs": 192, + 'dropout_rate': 0.1, + 'learning_rate': 0.0025, + 'one_minus_beta1': 0.1, + 'beta2': 0.9955159689799007, + 'weight_decay': 0.08121616522670176, + 'warmup_factor': 0.02, + 'weight_lr_power': 2, + 'label_smoothing': 0.2, + 'r': 0.75, + 'conformer_bs': 192, } class AdamWScheduleFree(torch.optim.Optimizer): - r"""Schedule Free AdamW + r"""Schedule Free AdamW""" + + def __init__( + self, + params, + lr=1e-3, + betas=(0.9, 0.999), + eps=1e-8, + weight_decay=0, + weight_lr_power=2, + warmup_steps=0, + r=0, + ): + defaults = dict( + lr=lr, + betas=betas, + eps=eps, + r=r, + k=0, + weight_sum=0.0, + lr_max=0.0, + warmup_steps=warmup_steps, + weight_lr_power=weight_lr_power, + weight_decay=weight_decay, + ) + + super().__init__(params, defaults) + + def reset(self): + for group in self.param_groups: + group['k'] = 0 + group['lr_max'] = 0 + group['weight_sum'] = 0 + + for p in group['params']: + # State initialization + state = self.state[p] + state['z'].copy_(state['x0']) + p.data.copy_(state['x0']) + state['exp_avg_sq'].zero_() + + def step(self, closure): + """Performs a single optimization step. + + Arguments: + closure (callable, optional): A closure that reevaluates the model + and returns the loss. """ - def __init__(self, params, - lr=1e-3, - betas=(0.9, 0.999), - eps=1e-8, - weight_decay=0, - weight_lr_power=2, - warmup_steps=0, - r=0, - ): - defaults = dict(lr=lr, - betas=betas, - eps=eps, - r=r, - k=0, - weight_sum=0.0, - lr_max=0.0, - warmup_steps=warmup_steps, - weight_lr_power=weight_lr_power, - weight_decay=weight_decay) - - super().__init__(params, defaults) - - def reset(self): - for group in self.param_groups: - group['k'] = 0 - group['lr_max'] = 0 - group['weight_sum'] = 0 - - for p in group['params']: - # State initialization - state = self.state[p] - state['z'].copy_(state['x0']) - p.data.copy_(state['x0']) - state['exp_avg_sq'].zero_() - - def step(self, closure): - """Performs a single optimization step. - - Arguments: - closure (callable, optional): A closure that reevaluates the model - and returns the loss. - """ - # Swap to extrapolated point: - for group in self.param_groups: - beta1, beta2 = group['betas'] - r = group['r'] - k = group['k'] - - for p in group['params']: - # State initialization - state = self.state[p] - if 'z' not in state: - state['z'] = torch.clone(p.data) - state['exp_avg_sq'] = torch.zeros_like(p.data, dtype=torch.bfloat16) - state['x0'] = p.data.cpu() - - z = state['z'] - - # Extrapolate - #p = p + (1-beta1)*(z-p) - #p.data.mul_(beta1).add_(z, alpha=1-beta1) - p.data.lerp_(end=z, weight=1-beta1) - - # Evaluate gradient at extrapolated point - loss = closure() - - for group in self.param_groups: - eps = group['eps'] - k = group['k'] - warmup_steps = group['warmup_steps'] - - if k < warmup_steps: - sched = (k+1) / warmup_steps - else: - sched = 1.0 - annealed_lr = group['lr']*sched - - lr = max(annealed_lr, eps) - - decay = group['weight_decay'] - beta1, beta2 = group['betas'] - weight_lr_power = group['weight_lr_power'] - - r = group['r'] - lr_max = group['lr_max'] = max(lr, group['lr_max']) - - weight = ((k+1)**r) * (lr_max**weight_lr_power) - weight_sum = group['weight_sum'] = group['weight_sum'] + weight - - ckp1 = weight/weight_sum - - bias_correction2 = 1 - beta2 ** (k+1) - step_size = lr * math.sqrt(bias_correction2) - - for p in group['params']: - if p.grad is None: - continue - grad = p.grad.data - - state = self.state[p] - - exp_avg_sq = state['exp_avg_sq'] - z = state['z'] - - # Unextrapolate - #p = (p - (1-beta1)*z)/beta1 - #p.data.sub_(z, alpha=1-beta1).div_(beta1) - p.data.lerp_(end=z, weight=1-1/beta1) - - # Decay the first and second moment running average coefficient - exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1-beta2) - denom = exp_avg_sq.sqrt().add_(eps) - - z.addcdiv_(grad, denom, value=-step_size) - - # Decay - z.sub_(p.data, alpha=step_size*decay) - - ### Take step - #p.data.mul_(1-ckp1).add_(z, alpha=ckp1) - p.data.lerp_(end=z, weight=ckp1) - - group['k'] = k+1 - return loss - -def init_optimizer_state(workload: spec.Workload, - model_params: spec.ParameterContainer, - model_state: spec.ModelAuxiliaryState, - hyperparameters: spec.Hyperparameters, - rng: spec.RandomState) -> spec.OptimizerState: + # Swap to extrapolated point: + for group in self.param_groups: + beta1, beta2 = group['betas'] + r = group['r'] + k = group['k'] + + for p in group['params']: + # State initialization + state = self.state[p] + if 'z' not in state: + state['z'] = torch.clone(p.data) + state['exp_avg_sq'] = torch.zeros_like(p.data, dtype=torch.bfloat16) + state['x0'] = p.data.cpu() + + z = state['z'] + + # Extrapolate + # p = p + (1-beta1)*(z-p) + # p.data.mul_(beta1).add_(z, alpha=1-beta1) + p.data.lerp_(end=z, weight=1 - beta1) + + # Evaluate gradient at extrapolated point + loss = closure() + + for group in self.param_groups: + eps = group['eps'] + k = group['k'] + warmup_steps = group['warmup_steps'] + + if k < warmup_steps: + sched = (k + 1) / warmup_steps + else: + sched = 1.0 + annealed_lr = group['lr'] * sched + + lr = max(annealed_lr, eps) + + decay = group['weight_decay'] + beta1, beta2 = group['betas'] + weight_lr_power = group['weight_lr_power'] + + r = group['r'] + lr_max = group['lr_max'] = max(lr, group['lr_max']) + + weight = ((k + 1) ** r) * (lr_max**weight_lr_power) + weight_sum = group['weight_sum'] = group['weight_sum'] + weight + + ckp1 = weight / weight_sum + + bias_correction2 = 1 - beta2 ** (k + 1) + step_size = lr * math.sqrt(bias_correction2) + + for p in group['params']: + if p.grad is None: + continue + grad = p.grad.data + + state = self.state[p] + + exp_avg_sq = state['exp_avg_sq'] + z = state['z'] + + # Unextrapolate + # p = (p - (1-beta1)*z)/beta1 + # p.data.sub_(z, alpha=1-beta1).div_(beta1) + p.data.lerp_(end=z, weight=1 - 1 / beta1) + + # Decay the first and second moment running average coefficient + exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2) + denom = exp_avg_sq.sqrt().add_(eps) + + z.addcdiv_(grad, denom, value=-step_size) + + # Decay + z.sub_(p.data, alpha=step_size * decay) + + ### Take step + # p.data.mul_(1-ckp1).add_(z, alpha=ckp1) + p.data.lerp_(end=z, weight=ckp1) + + group['k'] = k + 1 + return loss + + +def init_optimizer_state( + workload: spec.Workload, + model_params: spec.ParameterContainer, + model_state: spec.ModelAuxiliaryState, + hyperparameters: spec.Hyperparameters, + rng: spec.RandomState, +) -> spec.OptimizerState: del model_state optimizer = AdamWScheduleFree( @@ -170,43 +177,55 @@ def init_optimizer_state(workload: spec.Workload, warmup_steps=int(HPARAMS['warmup_factor'] * workload.step_hint * 0.75), weight_decay=HPARAMS['weight_decay'], weight_lr_power=HPARAMS['weight_lr_power'], - r=HPARAMS['r']) - - optimizer_state = {'optimizer':optimizer, 'max_checked_eval_step': -1, 'has_forced_reset': False, 'first_eval': False, } + r=HPARAMS['r'], + ) + + optimizer_state = { + 'optimizer': optimizer, + 'max_checked_eval_step': -1, + 'has_forced_reset': False, + 'first_eval': False, + } return optimizer_state -def update_params(workload: spec.Workload, - current_param_container: spec.ParameterContainer, - current_params_types: spec.ParameterTypeTree, - model_state: spec.ModelAuxiliaryState, - hyperparameters: spec.Hyperparameters, - batch: Dict[str, spec.Tensor], - loss_type: spec.LossType, - optimizer_state: spec.OptimizerState, - eval_results: List[Tuple[int, float]], - global_step: int, - rng: spec.RandomState) -> spec.UpdateReturn: + +def update_params( + workload: spec.Workload, + current_param_container: spec.ParameterContainer, + current_params_types: spec.ParameterTypeTree, + model_state: spec.ModelAuxiliaryState, + hyperparameters: spec.Hyperparameters, + batch: Dict[str, spec.Tensor], + loss_type: spec.LossType, + optimizer_state: spec.OptimizerState, + eval_results: List[Tuple[int, float]], + global_step: int, + rng: spec.RandomState, +) -> spec.UpdateReturn: """Return (updated_optimizer_state, updated_params, updated_model_state).""" del current_params_types del loss_type del hyperparameters - metric_name = workload.target_metric_name # TODO - remove force_reset eval_step = len(eval_results) - if (global_step > workload.step_hint*0.10) and optimizer_state['max_checked_eval_step'] < eval_step: + if (global_step > workload.step_hint * 0.10) and optimizer_state[ + 'max_checked_eval_step' + ] < eval_step: optimizer_state['max_checked_eval_step'] = eval_step # Don't do resetting on workloads that don't run eval often enough - if len(eval_results) >= 4: # and optimizer_state["first_eval_is_far_from_target"]: - val_metric = f"validation/{metric_name}" + if ( + len(eval_results) >= 4 + ): # and optimizer_state["first_eval_is_far_from_target"]: + val_metric = f'validation/{metric_name}' initial_eval = eval_results[0][1][val_metric] latest_eval = eval_results[-1][1][val_metric] second_latest_eval = eval_results[-2][1][val_metric] third_latest_eval = eval_results[-3][1][val_metric] fourth_latest_eval = eval_results[-4][1][val_metric] MARGIN = 0.01 - if metric_name in ["loss", "wer"]: + if metric_name in ['loss', 'wer']: # Decreasing eval workloads should be flipped initial_eval = -initial_eval latest_eval = -latest_eval @@ -216,10 +235,15 @@ def update_params(workload: spec.Workload, # Higher is better # scale as a curve from 0 --> 1 # if the eval values are far from the target (i.e. - worse than initial) and stays far from the target for 4 evals - if (latest_eval - initial_eval < MARGIN) and (latest_eval - second_latest_eval < MARGIN) and (second_latest_eval - third_latest_eval < MARGIN) and (third_latest_eval - fourth_latest_eval < MARGIN): + if ( + (latest_eval - initial_eval < MARGIN) + and (latest_eval - second_latest_eval < MARGIN) + and (second_latest_eval - third_latest_eval < MARGIN) + and (third_latest_eval - fourth_latest_eval < MARGIN) + ): # Reset parameters since we appear to have diverged - logging.info("Reseting All Weights ") - logging.info(f"Global Step: {global_step}") + logging.info('Reseting All Weights ') + logging.info(f'Global Step: {global_step}') optimizer_state['has_forced_reset'] = True # Perform reset @@ -229,7 +253,7 @@ def update_params(workload: spec.Workload, # Decrease learning rate by 2x if it diverged. for param_group in optimizer_state['optimizer'].param_groups: - param_group['lr'] = param_group['lr']/2.0 + param_group['lr'] = param_group['lr'] / 2.0 ########### @@ -243,18 +267,20 @@ def closure(): optimizer_state['optimizer'].zero_grad() logits_batch, new_model_state = workload.model_fn( - params=current_model, - augmented_and_preprocessed_input_batch=batch, - model_state=model_state, - mode=spec.ForwardPassMode.TRAIN, - rng=rng, - update_batch_norm=True) + params=current_model, + augmented_and_preprocessed_input_batch=batch, + model_state=model_state, + mode=spec.ForwardPassMode.TRAIN, + rng=rng, + update_batch_norm=True, + ) loss_dict = workload.loss_fn( - label_batch=batch['targets'], - logits_batch=logits_batch, - mask_batch=batch.get('weights'), - label_smoothing=HPARAMS['label_smoothing']) + label_batch=batch['targets'], + logits_batch=logits_batch, + mask_batch=batch.get('weights'), + label_smoothing=HPARAMS['label_smoothing'], + ) summed_loss = loss_dict['summed'] n_valid_examples = loss_dict['n_valid_examples'] if USE_PYTORCH_DDP: @@ -276,7 +302,7 @@ def get_batch_size(workload_name): if workload_name == 'criteo1tb': return 262_144 elif workload_name == 'fastmri': - return 16 # 32 + return 16 # 32 elif workload_name == 'imagenet_resnet': return 1024 elif workload_name == 'imagenet_vit': @@ -284,7 +310,7 @@ def get_batch_size(workload_name): elif workload_name == 'librispeech_conformer': return 224 elif workload_name == 'librispeech_deepspeech': - return 128 # 256 + return 128 # 256 elif workload_name == 'ogbg': return 512 elif workload_name == 'wmt': @@ -298,14 +324,17 @@ def get_batch_size(workload_name): else: raise ValueError(f'Unsupported workload name: {workload_name}.') -def data_selection(workload: spec.Workload, - input_queue: Iterator[Dict[str, spec.Tensor]], - optimizer_state: spec.OptimizerState, - current_param_container: spec.ParameterContainer, - model_state: spec.ModelAuxiliaryState, - hyperparameters: spec.Hyperparameters, - global_step: int, - rng: spec.RandomState) -> Dict[str, spec.Tensor]: + +def data_selection( + workload: spec.Workload, + input_queue: Iterator[Dict[str, spec.Tensor]], + optimizer_state: spec.OptimizerState, + current_param_container: spec.ParameterContainer, + model_state: spec.ModelAuxiliaryState, + hyperparameters: spec.Hyperparameters, + global_step: int, + rng: spec.RandomState, +) -> Dict[str, spec.Tensor]: """Select data from the infinitely repeating, pre-shuffled input queue. Each element of the queue is a batch of training examples and labels. """ @@ -318,4 +347,3 @@ def data_selection(workload: spec.Workload, del rng batch = next(input_queue) return batch - From 023a5c533cb51f95853a1ffa2693ae9bbfa27137 Mon Sep 17 00:00:00 2001 From: init-22 Date: Thu, 14 Aug 2025 18:04:42 +0000 Subject: [PATCH 7/7] removing unnecessary space --- README.md | 1 - 1 file changed, 1 deletion(-) diff --git a/README.md b/README.md index 525dc9210..0666d21d5 100644 --- a/README.md +++ b/README.md @@ -37,7 +37,6 @@ Submissions are evaluated based on their "time-to-result", i.e., the wall-clock --- > [!IMPORTANT] - > For future iterations of the AlgoPerf: Training Algorithms benchmark competition, we are switching to a rolling leaderboard, making a few changes to the competition rules, and also run all selected submissions on our hardware. **To submit your algorithm to the next iteration of the benchmark, please see our [How to Submit](#how-to-submit) section and the [submission repository](https://github.com/mlcommons/submissions_algorithms) which hosts the up to date AlgoPerf leaderboard.** ## Table of Contents