{"record":{"id":"c8439e443ab28b33","repo":"huggingface/pytorch-image-models","slug":"weight-decay-option-is-not-compatible-with-sparse","errorCode":null,"errorMessage":"weight_decay option is not compatible with sparse gradients","messagePattern":"weight_decay option is not compatible with sparse gradients","errorType":"exception","errorClass":"RuntimeError","httpStatus":null,"severity":"error","filePath":"timm/optim/madgrad.py","lineNumber":135,"sourceCode":"                if len(state) == 0:\n                    state['step'] = 0\n                    state['grad_sum_sq'] = torch.zeros_like(p)\n                    state['s'] = torch.zeros_like(p)\n                    if momentum != 0:\n                        state['x0'] = torch.clone(p).detach()\n\n                state['step'] += 1\n                grad_sum_sq = state['grad_sum_sq']\n                s = state['s']\n                lamb = lr * math.sqrt(state['step'])\n\n                # Apply weight decay\n                if weight_decay != 0:\n                    if group['decoupled_decay']:\n                        p.mul_(1.0 - group['lr'] * weight_decay)\n                    else:\n                        if grad.is_sparse:\n                            raise RuntimeError(\"weight_decay option is not compatible with sparse gradients\")\n                        grad.add_(p, alpha=weight_decay)\n\n                if grad.is_sparse:\n                    grad = grad.coalesce()\n                    grad_val = grad._values()\n\n                    p_masked = p.sparse_mask(grad)\n                    grad_sum_sq_masked = grad_sum_sq.sparse_mask(grad)\n                    s_masked = s.sparse_mask(grad)\n\n                    # Compute x_0 from other known quantities\n                    rms_masked_vals = grad_sum_sq_masked._values().pow(1 / 3).add_(eps)\n                    x0_masked_vals = p_masked._values().addcdiv(s_masked._values(), rms_masked_vals, value=1)\n\n                    # Dense + sparse op\n                    grad_sq = grad * grad\n                    grad_sum_sq.add_(grad_sq, alpha=lamb)\n                    grad_sum_sq_masked.add_(grad_sq, alpha=lamb)","sourceCodeStart":117,"sourceCodeEnd":153,"githubUrl":"https://github.com/huggingface/pytorch-image-models/blob/9a5261e31b3b5128526eb2658333b4c0a54464ae/timm/optim/madgrad.py#L117-L153","documentation":"In step(), MADGRAD with coupled (non-decoupled) weight decay adds weight_decay * p into the gradient, which cannot be done on a sparse tensor, so sparse gradients plus nonzero coupled weight decay raise a RuntimeError.","triggerScenarios":"A param group with weight_decay != 0 and decoupled_decay=False whose parameter receives a sparse gradient, when step() runs.","commonSituations":"Sparse embedding training where the config also enables L2 weight decay; default MADGRAD settings use coupled decay.","solutions":["Set decoupled_decay=True for the sparse group (decoupled decay multiplies params directly and works with sparse grads)","Set weight_decay=0 for the sparse param group","Make gradients dense or use a sparse-capable optimizer"],"exampleFix":"# before\nopt = MADGRAD([{'params': emb.parameters()}], lr=1e-3, weight_decay=1e-4)\n# after\nopt = MADGRAD([{'params': emb.parameters(), 'decoupled_decay': True}], lr=1e-3, weight_decay=1e-4)","handlingStrategy":"validation","validationCode":"assert decoupled_decay or weight_decay == 0 or not any(\n    p.grad is not None and p.grad.is_sparse for p in sparse_params)","typeGuard":null,"tryCatchPattern":"try:\n    opt.step()\nexcept RuntimeError as e:\n    if 'weight_decay' in str(e) and 'sparse' in str(e):\n        for g in opt.param_groups: g['decoupled_decay'] = True\n    else:\n        raise","preventionTips":["Prefer decoupled_decay=True when training with sparse embeddings","Keep weight_decay=0 on sparse groups"],"tags":["optimizer","madgrad","sparse-gradients","weight-decay"],"backgroundTag":"sparse-gradient-unsupported","analyzedSha":"9a5261e31b3b5128526eb2658333b4c0a54464ae","analyzedAt":"2026-08-27T02:34:25.417Z","schemaVersion":2},"datasetVersion":"2026-08-27T03:17:27.898Z"}