Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
120 changes: 114 additions & 6 deletions moe_infinity/runtime/model_offload.py
Original file line number Diff line number Diff line change
Expand Up @@ -165,6 +165,107 @@ def _out_block(experts_prefix: str) -> str:
del state_dict[down_key]


_GPT_OSS_EXPERT_FIELDS = (
"gate_up_proj_blocks",
"gate_up_proj_scales",
"gate_up_proj_bias",
"down_proj_blocks",
"down_proj_scales",
"down_proj_bias",
)


def _expand_gpt_oss_packed_experts(state_dict, config):
if getattr(config, "model_type", "") != "gpt_oss":
return

layer_prefixes = sorted(
{
key.rsplit(".", 1)[0]
for key in state_dict
if key.endswith(".mlp.experts.gate_up_proj_blocks")
}
)
expected_experts = int(config.num_local_experts)
for prefix in layer_prefixes:
packed = {}
missing = []
for field in _GPT_OSS_EXPERT_FIELDS:
key = f"{prefix}.{field}"
if key not in state_dict:
missing.append(field)
else:
packed[field] = state_dict[key]
if missing:
raise ValueError(
f"Incomplete GPT-OSS packed expert layer {prefix}: {missing}"
)
for field, tensor in packed.items():
if tensor.shape[0] != expected_experts:
raise ValueError(
f"{prefix}.{field} has {tensor.shape[0]} experts; "
f"expected {expected_experts}"
)
for expert_idx in range(expected_experts):
expert_prefix = f"{prefix}.{expert_idx}"
for field in _GPT_OSS_EXPERT_FIELDS:
view = packed[field][expert_idx]
if not view.is_contiguous():
raise ValueError(
f"Non-contiguous GPT-OSS slice {expert_prefix}.{field}"
)
state_dict[f"{expert_prefix}.{field}"] = view
for field in _GPT_OSS_EXPERT_FIELDS:
del state_dict[f"{prefix}.{field}"]


def _gpt_oss_expert_groups(name_id_map, config):
fields = {name: index for index, name in enumerate(_GPT_OSS_EXPERT_FIELDS)}
grouped = {}
for name, tensor_id in name_id_map.items():
layer_id, expert_id = parse_expert_id(name, config)
if layer_id is None or expert_id is None:
continue
field = name.rsplit(".", 1)[-1]
if field not in fields:
continue
slots = grouped.setdefault(
(layer_id, expert_id), [None] * len(_GPT_OSS_EXPERT_FIELDS)
)
slots[fields[field]] = tensor_id

topology = []
for layer_id in range(config.num_hidden_layers):
experts = []
for expert_id in range(config.num_local_experts):
ids = grouped.get((layer_id, expert_id))
if ids is None or any(tensor_id is None for tensor_id in ids):
raise ValueError(
f"Missing GPT-OSS expert tensors for layer={layer_id}, "
f"expert={expert_id}"
)
experts.append(ids)
topology.append((f"model.layers.{layer_id}.mlp.experts", experts))
return topology


def _make_expert_tensor_map(name_id_map, config):
if getattr(config, "model_type", "") == "gpt_oss":
return {
(layer_id, expert_id): tensor_ids[0]
for layer_id, (_, experts) in enumerate(
_gpt_oss_expert_groups(name_id_map, config)
)
for expert_id, tensor_ids in enumerate(experts)
}
result = {}
for name, tensor_id in name_id_map.items():
layer_id, expert_id = parse_expert_id(name, config)
if expert_id is not None:
result[(layer_id, expert_id)] = tensor_id
return result


def _identify_fp8_blockwise_pairs(keys):
key_set = set(keys)
pairs = []
Expand Down Expand Up @@ -707,6 +808,7 @@ def archer_from_pretrained(cls, *args, **kwargs):
state_dict = torch.load(ckpt)

_remap_v5_batched_experts(state_dict, self.config)
_expand_gpt_oss_packed_experts(state_dict, self.config)

is_gptq_ckpt = is_gptq_quantized(self.config)
_arch0_cast = (
Expand All @@ -725,6 +827,7 @@ def archer_from_pretrained(cls, *args, **kwargs):

if (
is_mxfp4_ckpt
and self.config.model_type != "gpt_oss"
and os.environ.get("MOE_INFINITY_MXFP4_DEQUANT", "")
== "1"
):
Expand Down Expand Up @@ -924,11 +1027,9 @@ def archer_from_pretrained(cls, *args, **kwargs):
self.name_id_map["embed_tokens.weight"] = 0
model.model.embed_tokens.weight.ar_id = 0

self.expert_tensor_map = dict()
for name, id in self.name_id_map.items():
layer_id, expert_id = parse_expert_id(name, self.config)
if expert_id is not None:
self.expert_tensor_map[(layer_id, expert_id)] = id
self.expert_tensor_map = _make_expert_tensor_map(
self.name_id_map, self.config
)
self.expert_prefetcher.expert_tensor_map = (
self.expert_tensor_map
)
Expand Down Expand Up @@ -1301,6 +1402,13 @@ def get_topology(self, model):
name_lst = []
ret_dict = {}

if getattr(self.config, "model_type", "") == "gpt_oss":
gpt_oss_topology = _gpt_oss_expert_groups(
self.name_id_map, self.config
)
name_lst.extend(name for name, _ in gpt_oss_topology)
ret_dict.update(dict(gpt_oss_topology))

for name, _ in model.named_parameters(recurse=True):
match = re.search(r"\d+", name)
if name not in self.name_id_map:
Expand Down Expand Up @@ -1509,7 +1617,7 @@ def gen_args_hook(
key = key.split(".")[0]
output_device_index = 0

if "expert" in key and self.config.model_type != "gpt_oss":
if "expert" in key:
for expert_idx, expert_tensors in enumerate(tensors):
expert_key = (
f"{key}.expert_{expert_idx}"
Expand Down
9 changes: 6 additions & 3 deletions moe_infinity/utils/hf_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -216,13 +216,16 @@ def parse_expert_id(
layer_id = int(layer_id)
expert_id = int(expert_id)
elif "gpt_oss" in arch or "gptoss" in arch:
layer_type = "decoder"
result = re.findall(
r"layers\.(\d+)\.mlp\.experts\.(gate_up_proj|down_proj)",
r"layers\.(\d+)\.mlp\.experts\.(\d+)\."
r"(?:gate_up_proj|down_proj)_(?:blocks|scales|bias)$",
param_name,
)
if result:
layer_id = int(result[0][0])
return layer_id, None
layer_id, expert_id = (int(value) for value in result[0])
if layer_id >= num_layers or expert_id >= config.num_local_experts:
return None, None

if result:
if layer_type == "decoder":
Expand Down
38 changes: 21 additions & 17 deletions tests/test_gpt_oss_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,37 +37,41 @@ def test_parse_moe_param_gpt_oss():
assert num_encoder_layers == 0


def test_parse_expert_id_gpt_oss_packed():
def test_parse_expert_id_gpt_oss_gate_up_slice():
from moe_infinity.utils.hf_config import parse_expert_id

config = make_gpt_oss_config()
layer_id, expert_id = parse_expert_id(
"model.layers.5.mlp.experts.gate_up_proj_blocks", config
)
assert layer_id == 5
assert expert_id is None
assert parse_expert_id(
"model.layers.5.mlp.experts.17.gate_up_proj_blocks", config
) == (5, 17)


def test_parse_expert_id_gpt_oss_router():
def test_parse_expert_id_gpt_oss_down_slice():
from moe_infinity.utils.hf_config import parse_expert_id

config = make_gpt_oss_config()
layer_id, expert_id = parse_expert_id(
"model.layers.5.mlp.router.weight", config
)
assert layer_id is None
assert expert_id is None
assert parse_expert_id(
"model.layers.11.mlp.experts.31.down_proj_bias", config
) == (11, 31)


def test_parse_expert_id_gpt_oss_down_proj():
def test_parse_expert_id_gpt_oss_rejects_out_of_range_expert():
from moe_infinity.utils.hf_config import parse_expert_id

config = make_gpt_oss_config()
assert parse_expert_id(
"model.layers.5.mlp.experts.32.gate_up_proj_blocks", config
) == (None, None)


def test_parse_expert_id_gpt_oss_router():
from moe_infinity.utils.hf_config import parse_expert_id

config = make_gpt_oss_config()
layer_id, expert_id = parse_expert_id(
"model.layers.11.mlp.experts.down_proj_blocks", config
assert parse_expert_id("model.layers.5.mlp.router.weight", config) == (
None,
None,
)
assert layer_id == 11
assert expert_id is None


def test_parse_expert_dtype_gpt_oss_none():
Expand Down
131 changes: 131 additions & 0 deletions tests/test_gpt_oss_offload_topology.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,131 @@
from types import SimpleNamespace

import torch

from moe_infinity.runtime.model_offload import (
_expand_gpt_oss_packed_experts,
_gpt_oss_expert_groups,
_make_expert_tensor_map,
)


def _config(layers=2, experts=128):
return SimpleNamespace(
architectures=["GptOssForCausalLM"],
model_type="gpt_oss",
num_hidden_layers=layers,
num_local_experts=experts,
)


def _packed_layer(layer_id, experts=128):
prefix = f"model.layers.{layer_id}.mlp.experts"
return {
f"{prefix}.gate_up_proj_blocks": torch.empty(
experts, 12, 8, dtype=torch.uint8
),
f"{prefix}.gate_up_proj_scales": torch.empty(
experts, 12, 1, dtype=torch.uint8
),
f"{prefix}.gate_up_proj_bias": torch.empty(
experts, 12, dtype=torch.bfloat16
),
f"{prefix}.down_proj_blocks": torch.empty(
experts, 6, 6, dtype=torch.uint8
),
f"{prefix}.down_proj_scales": torch.empty(
experts, 6, 1, dtype=torch.uint8
),
f"{prefix}.down_proj_bias": torch.empty(
experts, 6, dtype=torch.bfloat16
),
}


def test_expansion_creates_128_identities_per_layer_without_copy():
state = {**_packed_layer(0), **_packed_layer(1)}
originals = dict(state)

_expand_gpt_oss_packed_experts(state, _config())

prefixes = {
key.rsplit(".", 1)[0] for key in state if ".mlp.experts." in key
}
assert len(prefixes) == 128 * 2
assert len(state) == 128 * 2 * 6

for layer_id in range(2):
packed_prefix = f"model.layers.{layer_id}.mlp.experts"
for expert_idx in range(128):
expert_prefix = f"{packed_prefix}.{expert_idx}"
for field in (
"gate_up_proj_blocks",
"gate_up_proj_scales",
"gate_up_proj_bias",
"down_proj_blocks",
"down_proj_scales",
"down_proj_bias",
):
view = state[f"{expert_prefix}.{field}"]
packed = originals[f"{packed_prefix}.{field}"]
assert view.is_contiguous()
assert (
view.untyped_storage().data_ptr()
== packed.untyped_storage().data_ptr()
)
assert view.storage_offset() == expert_idx * packed.stride(0)


def test_expansion_rejects_incomplete_layer():
state = _packed_layer(0)
del state["model.layers.0.mlp.experts.down_proj_scales"]

try:
_expand_gpt_oss_packed_experts(state, _config(layers=1))
except ValueError as exc:
assert "down_proj_scales" in str(exc)
else:
raise AssertionError("incomplete GPT-OSS packed layer was accepted")


def _synthetic_name_id_map(layers=2, experts=128):
mapping = {}
tensor_id = 100
for layer_id in range(layers):
for expert_idx in range(experts):
prefix = f"model.layers.{layer_id}.mlp.experts.{expert_idx}"
for field in (
"gate_up_proj_blocks",
"gate_up_proj_scales",
"gate_up_proj_bias",
"down_proj_blocks",
"down_proj_scales",
"down_proj_bias",
):
mapping[f"{prefix}.{field}"] = tensor_id
tensor_id += 1
return mapping


def test_gpt_oss_topology_has_six_ordered_ids_for_every_expert():
config = _config()
name_id_map = _synthetic_name_id_map()

groups = _gpt_oss_expert_groups(name_id_map, config)

assert len(groups) == 2
assert all(len(experts) == 128 for _, experts in groups)
assert all(len(ids) == 6 for _, experts in groups for ids in experts)
first_ids = groups[0][1][0]
assert first_ids == [100, 101, 102, 103, 104, 105]


def test_expert_tensor_map_has_128_times_num_layers_entries():
config = _config()
name_id_map = _synthetic_name_id_map()

tensor_map = _make_expert_tensor_map(name_id_map, config)

assert len(tensor_map) == 128 * config.num_hidden_layers
assert tensor_map[(0, 0)] == 100
assert tensor_map[(1, 127)] == max(name_id_map.values()) - 5
Loading