from typing import Callable, List, Optional, Tuple, Union
import mlx.core as mx
from mlx.nn import Module
from mlx.utils import tree_flatten, tree_map, tree_merge, tree_reduce, tree_unflatten
class Optimizer:
def __init__(self, schedulers=None):
self._initialized = False
self._state = {"step": mx.array(0, mx.uint64)}
self._schedulers = {k: v for k, v in (schedulers or {}).items()}
def update(self, model: Module, gradients: dict):
model.update(self.apply_gradients(gradients, model))
def init(self, parameters: dict):
def update_state(params, state):
if isinstance(params, (list, tuple)):
state = list(state)
for i in range(len(state)):
state[i] = update_state(params[i], state[i])
if len(state) != len(params):
state.extend(tree_map(lambda _: {}, params[len(state) :]))
return type(params)(state)
elif isinstance(params, dict):
for k, v in params.items():
if k not in state:
state[k] = tree_map(lambda _: {}, v)
else:
state[k] = update_state(v, state[k])
return state
else:
return state
update_state(parameters, self._state)
tree_map(lambda p, s: s or self.init_single(p, s), parameters, self._state)
self._initialized = True
def init_single(self, parameter: mx.array, state: dict):
raise NotImplementedError()
def apply_gradients(self, gradients: dict, parameters: dict):
if not self._initialized:
self.init(gradients)
for param, scheduler in self._schedulers.items():
self.state[param] = scheduler(self.step)
self.state["step"] = self.step + 1
return tree_map(self.apply_single, gradients, parameters, self.state)
def apply_single(self, gradient: mx.array, parameter: mx.array, state: dict):
raise NotImplementedError()
@property
def state(self):
return self._state
@state.setter
def state(self, state: dict):
self._initialized = False
self._state = state
@property
def step(self):
return self.state["step"]
@property
def learning_rate(self):
return self.state["learning_rate"]
@learning_rate.setter
def learning_rate(self, learning_rate: Union[float, mx.array]):
self.state["learning_rate"] = mx.array(learning_rate)
def _maybe_schedule(
self, name: str, param: Union[float, Callable[[mx.array], mx.array]]
):
if isinstance(param, Callable):
self._schedulers[name] = param
parameter = param(self.step)
else:
parameter = mx.array(param)
self.state[name] = parameter
class MultiOptimizer(Optimizer):
def __init__(self, optimizers, filters: list = []):
super().__init__()
self._state = {}
if len(filters) != len(optimizers) - 1:
raise ValueError(
f"Given {len(filters)} filters but {len(optimizers)-1} needed."
)
self.optimizers = optimizers
self.filters = filters + [lambda *args, **kwargs: True]
def _split_dictionary(self, gradients: dict):
if len(self.optimizers) == 1:
return [gradients]
parts = [[] for _ in range(len(self.optimizers))]
flat_gradients = tree_flatten(gradients)
for k, g in flat_gradients:
for i, fn in enumerate(self.filters):
if fn(k, g):
parts[i].append((k, g))
break
return [tree_unflatten(p) for p in parts]
def init(self, parameters: dict):
for o, p in zip(self.optimizers, self._split_dictionary(parameters)):
o.init(p)
def apply_gradients(self, gradients: dict, parameters: dict):
tree = {}
for o, g in zip(self.optimizers, self._split_dictionary(gradients)):
tree = tree_merge(tree, o.apply_gradients(g, parameters))
return tree
@property
def state(self):
return {"states": [o.state for o in self.optimizers]}
@state.setter
def state(self, state: dict):
if "states" not in state or len(state["states"]) != len(self.optimizers):
raise ValueError("Invalid state provided")
for o, s in zip(self.optimizers, state["states"]):
o.state = s
@property
def learning_rate(self):
return self.optimizers[0].learning_rate
@learning_rate.setter
def learning_rate(self, learning_rate: Union[float, mx.array]):
for o in self.optimizers:
o.learning_rate = learning_rate
class SGD(Optimizer):
def __init__(
self,
learning_rate: Union[float, Callable[[mx.array], mx.array]],
momentum: float = 0.0,
weight_decay: float = 0.0,
dampening: float = 0.0,
nesterov: bool = False,
):
if nesterov and (momentum <= 0 or dampening != 0):
raise ValueError(
"Nesterov momentum requires a momentum and zero dampening."
)
super().__init__()
self._maybe_schedule("learning_rate", learning_rate)
self.momentum = momentum
self.weight_decay = weight_decay
self.dampening = dampening
self.nesterov = nesterov
def init_single(self, parameter: mx.array, state: dict):
state["v"] = mx.zeros_like(parameter)
def apply_single(self, gradient: mx.array, parameter: mx.array, state: dict):
if self.weight_decay != 0:
gradient += self.weight_decay * parameter
if self.momentum <= 0:
return parameter - self.learning_rate.astype(gradient.dtype) * gradient
v = self.momentum * state.get("v")
if self.dampening > 0:
v += (1 - self.dampening) * gradient
else:
v += gradient
if self.nesterov:
update = gradient + self.momentum * v
else:
update = v
state["v"] = v
return parameter - self.learning_rate.astype(gradient.dtype) * update
class RMSprop(Optimizer):
def __init__(
self,
learning_rate: Union[float, Callable[[mx.array], mx.array]],
alpha: float = 0.99,
eps: float = 1e-8,
):
super().__init__()
self._maybe_schedule("learning_rate", learning_rate)
self.alpha = alpha
self.eps = eps
if self.alpha < 0.0:
raise ValueError(
f"RMSprop alpha should be >=0, {self.alpha} was provided instead"
)
if self.eps < 0.0:
raise ValueError(
f"RMSprop epsilon should be >0, {self.eps} was provided instead"
)
def init_single(self, parameter: mx.array, state: dict):
state["v"] = mx.zeros_like(parameter)
def apply_single(self, gradient: mx.array, parameter: mx.array, state: dict):
lr = self.learning_rate.astype(gradient.dtype)
alpha = self.alpha
eps = self.eps
v = state["v"]
v = alpha * v + (1 - alpha) * mx.square(gradient)
state["v"] = v
return parameter - lr * gradient / (mx.sqrt(v) + eps)
class Adagrad(Optimizer):
def __init__(
self,
learning_rate: Union[float, Callable[[mx.array], mx.array]],
eps: float = 1e-8,
):
super().__init__()
self._maybe_schedule("learning_rate", learning_rate)
self.eps = eps
if self.eps < 0.0:
raise ValueError(
f"Adagrad epsilon should be >0, {self.eps} was provided instead"
)
def init_single(self, parameter: mx.array, state: dict):
state["v"] = mx.zeros_like(parameter)
def apply_single(self, gradient: mx.array, parameter: mx.array, state: dict):
lr = self.learning_rate.astype(gradient.dtype)
eps = self.eps
v = state["v"] + mx.square(gradient)
state["v"] = v
return parameter - lr * gradient / (mx.sqrt(v) + eps)
class AdaDelta(Optimizer):
def __init__(
self,
learning_rate: Union[float, Callable[[mx.array], mx.array]],
rho: float = 0.9,
eps: float = 1e-6,
):
super().__init__()
self._maybe_schedule("learning_rate", learning_rate)
self.rho = rho
self.eps = eps
if self.rho < 0.0:
raise ValueError(
f"AdaDelta rho should be >=0, {self.rho} was provided instead"
)
if self.eps < 0.0:
raise ValueError(
f"AdaDelta epsilon should be >0, {self.eps} was provided instead"
)
def init_single(self, parameter: mx.array, state: dict):
state["v"] = mx.zeros_like(parameter)
state["u"] = mx.zeros_like(parameter)
def apply_single(self, gradient: mx.array, parameter: mx.array, state: dict):
lr = self.learning_rate.astype(gradient.dtype)
rho = self.rho
eps = self.eps
v = state["v"]
u = state["u"]
v = rho * v + (1 - rho) * mx.square(gradient)
d = mx.sqrt(u + eps) / mx.sqrt(v + eps) * gradient
u = rho * u + (1 - rho) * mx.square(d)
state["v"] = v
state["u"] = u
return parameter - lr * d
class Adam(Optimizer):
def __init__(
self,
learning_rate: Union[float, Callable[[mx.array], mx.array]],
betas: List[float] = [0.9, 0.999],
eps: float = 1e-8,
bias_correction: bool = False,
):
super().__init__()
self._maybe_schedule("learning_rate", learning_rate)
self.betas = betas
self.eps = eps
self.bias_correction = bias_correction
def init_single(self, parameter: mx.array, state: dict):
state["m"] = mx.zeros_like(parameter)
state["v"] = mx.zeros_like(parameter)
def apply_single(self, gradient: mx.array, parameter: mx.array, state: dict):
lr = self.learning_rate.astype(gradient.dtype)
b1, b2 = self.betas
eps = self.eps
bias_correction = self.bias_correction
step = self.step
m = state["m"]
v = state["v"]
m = b1 * m + (1 - b1) * gradient
v = b2 * v + (1 - b2) * mx.square(gradient)
state["m"] = m
state["v"] = v
if bias_correction:
c1 = (lr / (1 - b1**step)).astype(gradient.dtype)
c2 = mx.rsqrt(1 - b2**step).astype(gradient.dtype)
numerator = c1 * m
denominator = mx.sqrt(v) * c2 + eps
return parameter - numerator / denominator
else:
return parameter - lr * m / (mx.sqrt(v) + eps)
class AdamW(Adam):
def __init__(
self,
learning_rate: Union[float, Callable[[mx.array], mx.array]],
betas: List[float] = [0.9, 0.999],
eps: float = 1e-8,
weight_decay: float = 0.01,
bias_correction: bool = False,
):
super().__init__(
learning_rate=learning_rate,
betas=betas,
eps=eps,
bias_correction=bias_correction,
)
self.weight_decay = weight_decay
def apply_single(self, gradient: mx.array, parameter: mx.array, state: dict):
lr = self.learning_rate.astype(gradient.dtype)
return super().apply_single(
gradient, parameter * (1 - lr * self.weight_decay), state
)
class Adamax(Adam):
def __init__(
self,
learning_rate: Union[float, Callable[[mx.array], mx.array]],
betas: List[float] = [0.9, 0.999],
eps: float = 1e-8,
):
super().__init__(learning_rate, betas, eps)
if not 0.0 <= eps:
raise ValueError(
f"Epsilon value should be >=0, {self.eps} was provided instead"
)
def init_single(self, parameter: mx.array, state: dict):
state["m"] = mx.zeros_like(parameter)
state["v"] = mx.zeros_like(parameter)
def apply_single(self, gradient: mx.array, parameter: mx.array, state: dict):
lr = self.learning_rate.astype(gradient.dtype)
b1, b2 = self.betas
eps = self.eps
m = state["m"]
v = state["v"]
m = b1 * m + (1 - b1) * gradient
v = mx.maximum(b2 * v, mx.abs(gradient))
state["m"] = m
state["v"] = v
return parameter - lr * m / (v + eps)
class Lion(Optimizer):
def __init__(
self,
learning_rate: Union[float, Callable[[mx.array], mx.array]],
betas: List[float] = [0.9, 0.99],
weight_decay: float = 0.0,
):
super().__init__()
self._maybe_schedule("learning_rate", learning_rate)
self.betas = betas
self.weight_decay = weight_decay
def init_single(self, parameter: mx.array, state: dict):
state["m"] = mx.zeros_like(parameter)
def apply_single(self, gradient: mx.array, parameter: mx.array, state: dict):
lr = self.learning_rate.astype(gradient.dtype)
b1, b2 = self.betas
weight_decay = self.weight_decay
m = state["m"]
c = b1 * m + (1 - b1) * gradient
state["m"] = b2 * m + (1 - b2) * gradient
if weight_decay > 0:
parameter = (1 - lr * weight_decay) * parameter
return parameter - lr * mx.sign(c)
class Adafactor(Optimizer):
def __init__(
self,
learning_rate: Union[float, Callable[[mx.array], mx.array], None] = None,
eps: Tuple[float, float] = (1e-30, 1e-3),
clip_threshold: float = 1.0,
decay_rate: float = -0.8,
beta_1: Optional[float] = None,
weight_decay: float = 0.0,
scale_parameter: bool = True,
relative_step: bool = True,
warmup_init: bool = False,
):
super().__init__()
if learning_rate is not None:
self._maybe_schedule("learning_rate", learning_rate)
self.eps = eps
self.clip_threshold = clip_threshold
self.decay_rate = decay_rate
self.beta_1 = beta_1
self.weight_decay = weight_decay
self.scale_parameter = scale_parameter
self.relative_step = relative_step
self.warmup_init = warmup_init
def init_single(self, parameter: mx.array, state: dict):
if parameter.ndim >= 2:
shape = parameter.shape
dtype = parameter.dtype
state["exp_avg_sq_row"] = mx.zeros(shape[:-1], dtype=dtype)
state["exp_avg_sq_col"] = mx.zeros(shape[:-2] + shape[-1:], dtype=dtype)
else:
state["exp_avg_sq"] = mx.zeros_like(parameter)
if self.beta_1 is not None:
state["exp_avg"] = mx.zeros_like(parameter)
def _compute_rms(self, inputs):
return mx.sqrt(mx.mean(mx.square(inputs)))
def _compute_learning_rate(self, step, parameter_rms):
if self.relative_step:
min_step = 1e-6 * step if self.warmup_init else 1e-2
relative_step_size = mx.minimum(min_step, mx.rsqrt(step))
else:
relative_step_size = self.learning_rate
relative_step_size = relative_step_size.astype(parameter_rms.dtype)
parameter_scale = 1.0
if self.scale_parameter:
parameter_scale = mx.maximum(self.eps[1], parameter_rms)
return parameter_scale * relative_step_size
def _approximate_exp_moving_avg(self, exp_avg_sq_row, exp_avg_sq_col):
r_factor = mx.rsqrt(
exp_avg_sq_row / mx.mean(exp_avg_sq_row, axis=-1, keepdims=True)
)
c_factor = mx.rsqrt(exp_avg_sq_col)
return mx.matmul(
mx.expand_dims(r_factor, axis=-1), mx.expand_dims(c_factor, axis=0)
)
def apply_single(self, gradient: mx.array, parameter: mx.array, state: dict):
factored = gradient.ndim >= 2
step = self.step
use_first_moment = self.beta_1 is not None
parameter_rms = self._compute_rms(parameter)
learning_rate = self._compute_learning_rate(step, parameter_rms)
beta_2 = 1.0 - (step**self.decay_rate).astype(parameter_rms.dtype)
update = mx.square(gradient) + self.eps[0]
if factored:
exp_avg_sq_row = state["exp_avg_sq_row"]
exp_avg_sq_col = state["exp_avg_sq_col"]
exp_avg_sq_row = (beta_2 * exp_avg_sq_row) + (
(1 - beta_2) * mx.mean(update, axis=-1)
)
exp_avg_sq_col = (beta_2 * exp_avg_sq_col) + (
(1 - beta_2) * mx.mean(update, axis=-2)
)
state["exp_avg_sq_row"] = exp_avg_sq_row
state["exp_avg_sq_col"] = exp_avg_sq_col
update = self._approximate_exp_moving_avg(exp_avg_sq_row, exp_avg_sq_col)
update = update * gradient
else:
exp_avg_sq = state["exp_avg_sq"]
exp_avg_sq = (beta_2 * exp_avg_sq) + ((1 - beta_2) * update)
state["exp_avg_sq"] = exp_avg_sq
update = mx.rsqrt(exp_avg_sq) * gradient
update = update / mx.maximum(
1.0, self._compute_rms(update) / self.clip_threshold
)
update = learning_rate * update
if use_first_moment:
exp_avg = state["exp_avg"]
exp_avg = (self.beta_1 * exp_avg) + ((1 - self.beta_1) * update)
state["exp_avg"] = exp_avg
update = exp_avg
if self.weight_decay != 0:
parameter += parameter * (-self.weight_decay * learning_rate)
return parameter - update
class Muon(Optimizer):
def __init__(
self,
learning_rate: Union[float, Callable[[mx.array], mx.array]],
momentum: float = 0.95,
weight_decay: float = 0.01,
nesterov: bool = True,
ns_steps: int = 5,
):
super().__init__()
self._maybe_schedule("learning_rate", learning_rate)
self.momentum = momentum
self.weight_decay = weight_decay
self.nesterov = nesterov
self.ns_steps = ns_steps
def init_single(self, parameter: mx.array, state: dict):
state["v"] = mx.zeros_like(parameter)
def _zeropower_via_newtonschulz5(self, X, steps: int):
assert (
X.ndim == 2
), f"Expected a 2D array for Newton-Schulz iteration, got shape {X.shape} instead."
a, b, c = (3.4445, -4.7750, 2.0315)
transpose_needed = X.shape[-2] > X.shape[-1]
if transpose_needed:
X = X.T
X = X / (mx.linalg.norm(X, keepdims=True) + 1e-7)
for _ in range(steps):
A = X @ X.T
B = mx.addmm(b * A, A, A, beta=1.0, alpha=c)
X = mx.addmm(a * X, B, X, beta=1.0, alpha=1.0)
if transpose_needed:
X = X.T
return X
def apply_single(self, gradient: mx.array, parameter: mx.array, state: dict):
if self.weight_decay != 0:
gradient = gradient + self.weight_decay * parameter
v = self.momentum * state["v"]
v = v + (1 - self.momentum) * gradient
state["v"] = v
if self.nesterov:
update = gradient * (1 - self.momentum) + v * self.momentum
else:
update = v
lr = self.learning_rate.astype(gradient.dtype)
if update.ndim >= 2:
original_shape = update.shape
reshape_needed = update.ndim > 2
if reshape_needed:
update = mx.reshape(update, (update.shape[0], -1))
update = self._zeropower_via_newtonschulz5(update, steps=self.ns_steps)
if reshape_needed:
update = mx.reshape(update, original_shape)
lr *= max(1, update.shape[-2] / update.shape[-1]) ** 0.5
return parameter - lr * update
def clip_grad_norm(grads, max_norm):
norm_squared = tree_reduce(lambda acc, g: acc + g.square().sum(), grads, 0.0)
total_norm = mx.sqrt(norm_squared)
normalizer = mx.minimum(max_norm / (total_norm + 1e-6), 1.0)
clipped_grads = tree_map(lambda g: g * normalizer, grads)
return clipped_grads, total_norm