| name | gradient-grouping-learning-rate-scaling |
| title | Taming LLMs by Scaling Learning Rates with Gradient Grouping |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2506.01049 |
| keywords | ["Optimization","Learning Rate Scheduling","Adaptive Methods","Training Stability"] |
| description | Improve adaptive learning rates by clustering gradient statistics within layers and applying cluster-specific scaling. |
SGG: Fine-Grained Learning Rate Control Without Manual Tuning
Training large language models is sensitive to learning rate: too high causes divergence, too low wastes computation. Most adaptive optimizers apply the same scaling globally across all parameters, missing structure in how different parts of the network train. SGG (Scaling with Gradient Grouping) clusters parameters within each layer into groups based on gradient statistics, then applies tailored scaling per group. This balances per-layer constraints with per-parameter precision, converging faster and more stably across varying batch sizes and learning rates.
Core Concept
Parameter-adaptive learning rates are common, but ignoring within-layer structure misses optimization opportunities. Parameters in a layer often split into distinct groups: some with consistently high gradients (critical decision points), others with low gradients (refinement parameters). SGG identifies these groups dynamically and scales each differently, preventing any single group from dominating updates while ensuring critical parameters stay responsive.
Architecture Overview
- Gradient Statistics Collection: Track gradient magnitude, variance within each layer
- Dynamic Clustering: Partition parameters into K clusters based on gradient statistics (typically K=3-5)
- Cluster-Specific Scaling: Compute separate scaling factors per cluster, normalizing gradient magnitudes
- Optimizer Wrapper: Integrates with existing optimizers (AdamW) as a lightweight post-processing step
- Stability Monitoring: Tracks gradient norm evolution to detect and prevent divergence
Implementation
This implementation demonstrates SGG as an optimizer wrapper for improved training stability.
Build the gradient grouping analyzer:
import torch
import torch.nn as nn
import numpy as np
from typing import Dict, List, Tuple
from collections import defaultdict
class GradientGroupingAnalyzer:
"""Analyze gradient statistics and identify parameter groups within layers."""
def __init__(self, num_clusters: int = 3, window_size: int = 100):
self.num_clusters = num_clusters
self.window_size = window_size
self.gradient_history = defaultdict(list)
def compute_gradient_statistics(self, model: nn.Module) -> Dict:
"""
Compute gradient magnitude and variance per layer.
Returns dict: layer_name -> {mean, std, percentiles}
"""
layer_stats = {}
for name, param in model.named_parameters():
if param.grad is None:
continue
grad = param.grad.detach()
grad_abs = torch.abs(grad)
stats = {
"mean": float(grad_abs.mean()),
"std": float(grad_abs.std()),
: (grad_abs.()),
: (grad_abs.()),
: (torch.quantile(grad_abs, )),
: (torch.quantile(grad_abs, )),
: (torch.quantile(grad_abs, )),
}
layer_name = .join(name.split()[:-])
layer_name layer_stats:
layer_stats[layer_name] = []
layer_stats[layer_name].append(stats)
layer_stats
() -> [, []]:
layer_params = {}
name, param model.named_parameters():
layer_name name param.grad :
grad = param.grad.detach().flatten()
layer_params[name] = grad
layer_params:
{}
features = []
param_names = (layer_params.keys())
name param_names:
grad = layer_params[name]
feature = [
(torch.(grad).mean()),
(torch.(grad).std()),
((grad ** ).mean()) **
]
features.append(feature)
features = np.array(features)
centroids = ._kmeans(features, .num_clusters)
clusters = ._assign_clusters(features, centroids)
cluster_mapping = defaultdict()
param_name, cluster_id (param_names, clusters):
cluster_mapping[cluster_id].append(param_name)
(cluster_mapping)
() -> np.ndarray:
indices = np.random.choice((data), k, replace=)
centroids = data[indices].copy()
iteration ():
distances = np.linalg.norm(data[:, ] - centroids, axis=)
assignments = np.argmin(distances, axis=)
new_centroids = np.array([
data[assignments == i].mean(axis=) (assignments == i).()
centroids[i]
i (k)
])
np.allclose(centroids, new_centroids):
centroids = new_centroids
centroids
() -> np.ndarray:
distances = np.linalg.norm(data[:, ] - centroids, axis=)
np.argmin(distances, axis=)
analyzer = GradientGroupingAnalyzer(num_clusters=)
model = nn.Sequential(
nn.Linear(, ),
nn.ReLU(),
nn.Linear(, )
)
dummy_input = torch.randn(, )
output = model(dummy_input).()
output.backward()
layer_stats = analyzer.compute_gradient_statistics(model)
layer, stats_list layer_stats.items():
()
clusters = analyzer.cluster_parameters(model, )
cluster_id, param_names clusters.items():
()
Implement SGG optimizer wrapper:
class ScalingWithGradientGrouping(torch.optim.Optimizer):
"""
Optimizer wrapper that applies cluster-specific gradient scaling.
Wraps AdamW or other standard optimizer.
"""
def __init__(self, model: nn.Module, optimizer_class=torch.optim.AdamW,
lr: float = 1e-4, num_clusters: int = 3, eps: float = 1e-8):
self.model = model
self.base_optimizer = optimizer_class(model.parameters(), lr=lr)
self.analyzer = GradientGroupingAnalyzer(num_clusters=num_clusters)
self.eps = eps
self.layer_clusters = {}
self.step_count = 0
def zero_grad(self):
"""Clear gradients."""
self.base_optimizer.zero_grad()
def step(self, closure=None):
"""
Optimizer step with gradient grouping and scaling.
"""
self.step_count += 1
if self.step_count % 100 == 0:
self._update_clusters()
self._apply_group_scaling()
return .base_optimizer.step(closure)
():
layer_names = ()
name, param .model.named_parameters():
param.grad :
layer_name = .join(name.split()[:-])
layer_names.add(layer_name)
layer_name layer_names:
.layer_clusters[layer_name] = \
.analyzer.cluster_parameters(.model, layer_name)
():
layer_name, clusters .layer_clusters.items():
cluster_id, param_names clusters.items():
cluster_grad_norm =
param_name param_names:
param = ._get_parameter_by_name(param_name)
param.grad :
cluster_grad_norm += torch.(param.grad ** )
cluster_grad_norm = torch.sqrt(cluster_grad_norm + .eps)
target_norm =
scale_factor = target_norm / (cluster_grad_norm + .eps)
param_name param_names:
param = ._get_parameter_by_name(param_name)
param.grad :
param.grad.mul_(scale_factor)
() -> torch.nn.Parameter:
parts = param_name.split()
obj = .model
part parts:
obj = (obj, part)
obj
model = nn.Sequential(
nn.Linear(, ),
nn.ReLU(),
nn.Linear(, )
)
optimizer = ScalingWithGradientGrouping(
model,
optimizer_class=torch.optim.AdamW,
lr=,
num_clusters=
)
()
step ():
optimizer.zero_grad()
x = torch.randn(, )
y = torch.randn(, )
output = model(x)
loss = nn.functional.mse_loss(output, y)
loss.backward()
optimizer.step()
(step + ) % == :
()
Benchmark SGG against standard AdamW:
def benchmark_optimizers():
"""Compare SGG vs standard AdamW convergence."""
torch.manual_seed(42)
model_sgg = nn.Sequential(nn.Linear(256, 512), nn.ReLU(), nn.Linear(512, 128))
model_adamw = nn.Sequential(nn.Linear(256, 512), nn.ReLU(), nn.Linear(512, 128))
with torch.no_grad():
for p1, p2 in zip(model_sgg.parameters(), model_adamw.parameters()):
p2.copy_(p1)
opt_sgg = ScalingWithGradientGrouping(model_sgg, lr=1e-3)
opt_adamw = torch.optim.AdamW(model_adamw.parameters(), lr=1e-3)
X = torch.randn(1000, 256)
y = torch.randn(1000, 128)
losses_sgg = []
losses_adamw = []
for epoch in range(50):
opt_sgg.zero_grad()
out_sgg = model_sgg(X)
loss_sgg = nn.functional.mse_loss(out_sgg, y)
loss_sgg.backward()
opt_sgg.step()
losses_sgg.append(loss_sgg.item())
opt_adamw.zero_grad()
out_adamw = model_adamw(X)
loss_adamw = nn.functional.mse_loss(out_adamw, y)
loss_adamw.backward()
opt_adamw.step()
losses_adamw.append(loss_adamw.item())
return losses_sgg, losses_adamw
losses_sgg, losses_adamw = benchmark_optimizers()
print(f"Final SGG loss: {losses_sgg[-]:f}")
()
()
Practical Guidance
| Aspect | Details |
|---|
| Number of Clusters | 3-5 typical; more clusters = finer control but higher overhead |
| Clustering Frequency | Update every 50-100 steps; balances adaptivity with compute |
| Learning Rate | Start with same LR as standard AdamW; may tolerate higher LR with SGG |
| Model Size | Benefits increase with model size; marginal gains on small models |
| Batch Size Sensitivity | SGG reduces LR tuning sensitivity across batch sizes |
When to Use:
- Training large language models on diverse batch sizes
- Sensitive to hyperparameter tuning; SGG reduces this burden
- Want faster convergence and better stability without architecture changes
- Compatible with other optimization techniques (gradient checkpointing, FSDP)
- Fine-tuning on multiple downstream tasks with single hyperparameter set
When NOT to Use:
- Simple models or datasets where standard AdamW already works well
- Real-time training with extremely tight latency constraints (clustering adds overhead)
- Already using highly specialized, hand-tuned learning rate schedules
- Theoretical analysis requires simple optimizer (stick with AdamW)
Common Pitfalls:
- Cluster count too high: wasted computation on fine-grained control with minimal benefit
- Clustering infrequently: outdated clusters become misaligned with actual gradient distribution
- Not accounting for warmup: clustering during warmup phase can be noisy; skip first N steps
- Ignoring divergence signals: monitor gradient norms; if still diverging, reduce learning rate globally
Reference
Taming LLMs by Scaling Learning Rates with Gradient Grouping
https://arxiv.org/abs/2506.01049