Views
No views yet
1import torch
2from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
3from kernels import get_kernel
4
5optimizer = get_kernel("motif-technologies/optimizer")
6get_default_muon_param_groups = optimizer.muon.get_default_muon_param_groups
7
8model = None # your model here
9fsdp_model = FSDP(model)
10
11# muon, in nature, cannot use 1-d tensor
12# we provide helper function to group such tensors
13# you can use your own function, if necessary
14params = get_default_muon_param_groups(model) # user can write own is_muon_func, if necessary
15
16optim = optimizer.Muon(
17 params,
18 lr=0.01,
19 momentum=0.9,
20 weight_decay=1e-4,
21)_StridedShard compatibility with PyTorch 2.10+.pip install pre-commit pre-commit install--style=file)pre-commit run --all-files pre-commit run isort --all-files