diff --git a/moe_infinity/runtime/model_offload.py b/moe_infinity/runtime/model_offload.py index 74a864d7..d1a7c393 100644 --- a/moe_infinity/runtime/model_offload.py +++ b/moe_infinity/runtime/model_offload.py @@ -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 = [] @@ -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 = ( @@ -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" ): @@ -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 ) @@ -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: @@ -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}" diff --git a/moe_infinity/utils/hf_config.py b/moe_infinity/utils/hf_config.py index b85b1853..ef0d4c85 100644 --- a/moe_infinity/utils/hf_config.py +++ b/moe_infinity/utils/hf_config.py @@ -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": diff --git a/tests/test_gpt_oss_config.py b/tests/test_gpt_oss_config.py index 5c90fef0..9ed5710e 100644 --- a/tests/test_gpt_oss_config.py +++ b/tests/test_gpt_oss_config.py @@ -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(): diff --git a/tests/test_gpt_oss_offload_topology.py b/tests/test_gpt_oss_offload_topology.py new file mode 100644 index 00000000..fdaa883b --- /dev/null +++ b/tests/test_gpt_oss_offload_topology.py @@ -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