[lora] More efficient pinned memory (#20876)

This commit is contained in:
Erik Wijmans
2026-05-30 09:04:59 +09:00
committed by GitHub
parent a5e6a8887a
commit 95cd2fd29f
5 changed files with 295 additions and 31 deletions
+3
View File
@@ -48,6 +48,7 @@ class LoRALayer(nn.Module):
# lora weights in cpu. The weights are loaded from checkpoint.
self.weights: Dict[str, torch.Tensor] = {}
self.pinned_weights: Dict[str, torch.Tensor] = {}
class LoRAAdapter(nn.Module):
@@ -87,7 +88,9 @@ class LoRAAdapter(nn.Module):
)
self.embedding_layers: Dict[str, torch.Tensor] = {}
self.pinned_embedding_layers: Dict[str, torch.Tensor] = {}
self.added_tokens_embeddings: Dict[str, torch.Tensor] = {}
self.pinned_added_tokens_embeddings: Dict[str, torch.Tensor] = {}
@staticmethod
def _build_moe_gated_map(base_model: torch.nn.Module) -> Dict[int, bool]:
+1 -4
View File
@@ -631,10 +631,6 @@ class LoRAManager:
)
lora_adapter.initialize_weights()
# If we want to overlap loading LoRA adapters with compute, they must be pinned in CPU memory
if self.enable_lora_overlap_loading:
lora_adapter.pin_weights_in_cpu()
self.loras[lora_ref.lora_id] = lora_adapter
def load_lora_weights_from_tensors(
@@ -707,6 +703,7 @@ class LoRAManager:
lora_added_tokens_size=self.lora_added_tokens_size,
experts_shared_outer_loras=self.experts_shared_outer_loras,
strict_loading=self.lora_strict_loading,
enable_lora_overlap_loading=self.enable_lora_overlap_loading,
)
# Initializing memory pool with base model
+223 -16
View File
@@ -1,6 +1,17 @@
import logging
import re
from typing import Callable, Dict, Iterable, Iterator, List, Optional, Set, Tuple, Union
from typing import (
Callable,
Dict,
Iterable,
Iterator,
List,
Optional,
Set,
Tuple,
Union,
overload,
)
import torch
@@ -22,12 +33,14 @@ from sglang.srt.lora.utils import (
REPLICATED_LINEAR_LORA_NAMES,
ROW_PARALLELISM_LINEAR_LORA_NAMES,
LoRAType,
copy_weight_into_buffer,
get_hidden_dim,
get_lm_head_lora_b_shard_size,
get_normalized_target_modules,
get_stacked_multiply,
get_target_module_name,
)
from sglang.srt.utils import is_pin_memory_available
from sglang.srt.utils.hf_transformers_utils import AutoConfig
logger = logging.getLogger(__name__)
@@ -53,6 +66,28 @@ class EmptySlot:
EMPTY_SLOT = EmptySlot()
@overload
def append_cache_key_suffix(cache_keys: str, suffix: str) -> str: ...
@overload
def append_cache_key_suffix(
cache_keys: Dict[int, str], suffix: str
) -> Dict[int, str]: ...
def append_cache_key_suffix(
cache_keys: Union[str, Dict[int, str]],
suffix: str,
) -> Union[str, Dict[int, str]]:
if isinstance(cache_keys, dict):
return {
expert_id: f"{cache_key}#{suffix}"
for expert_id, cache_key in cache_keys.items()
}
return f"{cache_keys}#{suffix}"
def _get_moe_ep_context() -> Tuple[int, int]:
"""Return `(moe_ep_size, moe_ep_rank)`, or `(1, 0)` if the MoE EP group
is not initialized (hermetic tests or pure-TP launches)."""
@@ -107,6 +142,7 @@ class LoRAMemoryPool:
lora_added_tokens_size: int,
experts_shared_outer_loras: bool = False,
strict_loading: bool = False,
enable_lora_overlap_loading: bool = False,
):
self.base_hf_config: AutoConfig = base_hf_config
self.num_layer: int = base_hf_config.num_hidden_layers
@@ -119,6 +155,8 @@ class LoRAMemoryPool:
self.target_modules: Set[str] = target_modules
self.experts_shared_outer_loras: bool = experts_shared_outer_loras
self.strict_loading: bool = strict_loading
self.enable_lora_overlap_loading: bool = enable_lora_overlap_loading
self.pin_memory_available: bool = is_pin_memory_available()
# Under EP with a Triton/DeepGEMM runner, `StandardDispatcher` remaps
# global `topk_ids` -> local expert IDs before the MoE kernel, so
@@ -258,18 +296,21 @@ class LoRAMemoryPool:
def _iter_local_expert_weights(
self,
weights: Union[torch.Tensor, Dict[int, torch.Tensor]],
) -> Iterator[Tuple[int, torch.Tensor]]:
"""Yield `(local_expert_id, weight)` pairs for per-expert MoE LoRA
inputs, filtered/remapped to this rank's slice. Accepts either a
`{global_eid: 2D tensor}` dict or a 3D `[num_experts, *, *]` tensor."""
cache_keys: Union[str, Dict[int, str]],
) -> Iterator[Tuple[int, torch.Tensor, str]]:
"""Yield `(local_expert_id, weight, cache_key)` triples for per-expert
MoE LoRA inputs, filtered/remapped to this rank's slice. Accepts either
a `{global_eid: 2D tensor}` dict or a 3D `[num_experts, *, *]` tensor."""
if isinstance(weights, dict):
assert isinstance(cache_keys, dict)
for gid, w in weights.items():
lid = self._global_to_local_expert_id(gid)
if lid is not None:
yield lid, w
yield lid, w, cache_keys[gid]
return
if isinstance(weights, torch.Tensor) and weights.dim() == 3:
assert isinstance(cache_keys, str)
total = weights.shape[0]
if self.moe_use_local_expert_ids:
start = self.moe_ep_rank * self._num_experts_local
@@ -277,7 +318,11 @@ class LoRAMemoryPool:
else:
start, count = 0, total
for i in range(count):
yield i, weights[start + i]
yield (
i,
weights[start + i],
append_cache_key_suffix(cache_keys, f"expert{start + i}"),
)
return
raise TypeError(
@@ -593,6 +638,35 @@ class LoRAMemoryPool:
self.get_lora_B_shape,
)
def _get_maybe_cached_weight_for_transfer(
self,
pinned_weight_store: Dict[str, torch.Tensor],
cache_key: str,
weight: torch.Tensor,
) -> torch.Tensor:
if (
not self.pin_memory_available
or weight.device.type != "cpu"
or weight.is_pinned()
):
return weight
if not self.enable_lora_overlap_loading:
return weight.pin_memory()
cached_weight = pinned_weight_store.get(cache_key)
if cached_weight is None:
cached_weight = weight.pin_memory()
pinned_weight_store[cache_key] = cached_weight
elif cached_weight.shape != weight.shape or cached_weight.dtype != weight.dtype:
raise ValueError(
f"LoRA pinned weight cache key collision for {cache_key!r}: "
f"cached shape={cached_weight.shape}, dtype={cached_weight.dtype}; "
f"new shape={weight.shape}, dtype={weight.dtype}."
)
return cached_weight
def prepare_lora_batch(
self,
cur_uids: Set[Optional[str]],
@@ -694,7 +768,7 @@ class LoRAMemoryPool:
assert (
buffer_view.shape == weight.shape
), f"LoRA buffer shape {buffer_view.shape} does not match weight shape {weight.shape}."
buffer_view.copy_(weight, non_blocking=True)
copy_weight_into_buffer(buffer_view, weight)
if uid is None:
for i in range(self.num_layer):
@@ -748,7 +822,9 @@ class LoRAMemoryPool:
logger.warning(msg)
for layer_id in range(self.num_layer):
layer_weights = lora_adapter.layers[layer_id].weights
layer = lora_adapter.layers[layer_id]
layer_weights = layer.weights
pinned_layer_weights = layer.pinned_weights
# - Standard: module_name -> torch.Tensor
# - MoE: module_name -> Dict[expert_id -> torch.Tensor]
temp_A_buffer: Dict[str, Union[torch.Tensor, Dict[int, torch.Tensor]]] = {
@@ -757,6 +833,12 @@ class LoRAMemoryPool:
temp_B_buffer: Dict[str, Union[torch.Tensor, Dict[int, torch.Tensor]]] = {
target_module: None for target_module in self.B_buffer
}
temp_A_cache_keys: Dict[str, Optional[Union[str, Dict[int, str]]]] = {
target_module: None for target_module in self.A_buffer
}
temp_B_cache_keys: Dict[str, Optional[Union[str, Dict[int, str]]]] = {
target_module: None for target_module in self.B_buffer
}
for name, weights in layer_weights.items():
target_module = get_target_module_name(name, self.target_modules)
@@ -770,25 +852,33 @@ class LoRAMemoryPool:
if temp_A_buffer[target_module] is None:
temp_A_buffer[target_module] = {}
temp_B_buffer[target_module] = {}
temp_A_cache_keys[target_module] = {}
temp_B_cache_keys[target_module] = {}
expert_id = int(expert_match.group(1))
if "lora_A" in name:
temp_A_buffer[target_module][expert_id] = weights
temp_A_cache_keys[target_module][expert_id] = name
else:
temp_B_buffer[target_module][expert_id] = weights
temp_B_cache_keys[target_module][expert_id] = name
elif "experts" in name and weights.dim() == 3:
# Shared outer MoE weight — 3D tensor [expert_dim, rank, hidden]
target_module = target_module + "_moe"
if "lora_A" in name:
temp_A_buffer[target_module] = weights
temp_A_cache_keys[target_module] = name
else:
temp_B_buffer[target_module] = weights
temp_B_cache_keys[target_module] = name
else:
# Standard weight — single tensor per module
if "lora_A" in name:
temp_A_buffer[target_module] = weights
temp_A_cache_keys[target_module] = name
else:
temp_B_buffer[target_module] = weights
temp_B_cache_keys[target_module] = name
# Track which buffer keys correspond to a real wrapped module on
# this layer. `temp_A/B_buffer` is seeded with every key in the
@@ -822,6 +912,12 @@ class LoRAMemoryPool:
target_module,
)
)
cache_keys = temp_A_cache_keys[target_module]
assert cache_keys is not None
temp_A_cache_keys[target_module] = append_cache_key_suffix(
cache_keys,
f"moe_tp{self.moe_tp_rank}",
)
if temp_B_buffer.get(target_module) is not None:
temp_B_buffer[target_module] = (
module.slice_moe_lora_b_weights(
@@ -830,6 +926,12 @@ class LoRAMemoryPool:
target_module,
)
)
cache_keys = temp_B_cache_keys[target_module]
assert cache_keys is not None
temp_B_cache_keys[target_module] = append_cache_key_suffix(
cache_keys,
f"moe_tp{self.moe_tp_rank}",
)
continue
@@ -849,9 +951,22 @@ class LoRAMemoryPool:
temp_A_buffer[target_module] = module.slice_lora_a_weights(
temp_A_buffer[target_module], self.tp_rank
)
cache_keys = temp_A_cache_keys[target_module]
assert cache_keys is not None
temp_A_cache_keys[target_module] = append_cache_key_suffix(
cache_keys,
f"tp{self.tp_rank}",
)
temp_B_buffer[target_module] = module.slice_lora_b_weights(
temp_B_buffer[target_module], self.tp_rank
)
cache_keys = temp_B_cache_keys[target_module]
assert cache_keys is not None
temp_B_cache_keys[target_module] = append_cache_key_suffix(
cache_keys,
f"tp{self.tp_rank}",
)
for name, weights in temp_A_buffer.items():
if name not in active_target_modules:
@@ -859,6 +974,7 @@ class LoRAMemoryPool:
c = get_stacked_multiply(name, self.base_model)
max_r = self.max_lora_rank
target_buffer = self.A_buffer[name][layer_id]
weights_cache_key = temp_A_cache_keys[name]
if name in ["gate_up_proj_moe", "down_proj_moe"]:
if self.experts_shared_outer_loras and name == "gate_up_proj_moe":
@@ -875,6 +991,12 @@ class LoRAMemoryPool:
f"gate_up_proj_moe lora_A has expert_dim="
f"{weights.shape[0]} (expected 1)."
)
assert isinstance(weights_cache_key, str)
weights = self._get_maybe_cached_weight_for_transfer(
pinned_layer_weights,
weights_cache_key,
weights,
)
representative_weight = weights[0]
buffer_view = target_buffer[
buffer_id, 0, : lora_rank * c, :
@@ -888,6 +1010,11 @@ class LoRAMemoryPool:
f"{len(weights)} entries (expected 1)."
)
rep = next(iter(weights.values()))
assert isinstance(weights_cache_key, dict)
rep_cache_key = next(iter(weights_cache_key.values()))
rep = self._get_maybe_cached_weight_for_transfer(
pinned_layer_weights, rep_cache_key, rep
)
representative_weight = rep
buffer_view = target_buffer[
buffer_id, 0, : lora_rank * c, :
@@ -921,11 +1048,21 @@ class LoRAMemoryPool:
# the MoE kernel's [:max_r] / [max_r:2*max_r] slicing
# is correct.
target_buffer[buffer_id].zero_()
for local_eid, expert_weight in self._iter_local_expert_weights(
weights
assert isinstance(weights_cache_key, (str, dict))
for (
local_eid,
expert_weight,
expert_cache_key,
) in self._iter_local_expert_weights(
weights, weights_cache_key
):
if expert_weight is None:
continue
expert_weight = self._get_maybe_cached_weight_for_transfer(
pinned_layer_weights,
expert_cache_key,
expert_weight,
)
for ci in range(c):
buffer_view = target_buffer[
buffer_id,
@@ -941,12 +1078,20 @@ class LoRAMemoryPool:
)
else:
buffer_view = target_buffer[buffer_id, : lora_rank * c, :]
if weights is not None:
assert isinstance(weights_cache_key, str)
weights = self._get_maybe_cached_weight_for_transfer(
pinned_layer_weights,
weights_cache_key,
weights,
)
load_lora_weight_tensor(buffer_view, weights)
for name, weights in temp_B_buffer.items():
if name not in active_target_modules:
continue
target_buffer = self.B_buffer[name][layer_id]
weights_cache_key = temp_B_cache_keys[name]
if name in ["gate_up_proj_moe", "down_proj_moe"]:
if self.experts_shared_outer_loras and name == "down_proj_moe":
@@ -962,11 +1107,17 @@ class LoRAMemoryPool:
)
buffer_view = target_buffer[buffer_id, 0, :, :lora_rank]
w = weights[0]
assert isinstance(weights_cache_key, str)
if w is not None:
w = w * lora_adapter.scaling
w = self._get_maybe_cached_weight_for_transfer(
pinned_layer_weights,
append_cache_key_suffix(
weights_cache_key, "expert0"
),
w,
)
load_lora_weight_tensor(buffer_view, w)
# Zero beyond loaded rank — MoE kernel reads full max_rank
target_buffer[buffer_id, 0, :, lora_rank:].zero_()
elif isinstance(weights, dict) and len(weights) > 0:
if len(weights) != 1:
raise ValueError(
@@ -975,37 +1126,65 @@ class LoRAMemoryPool:
f"{len(weights)} entries (expected 1)."
)
rep = next(iter(weights.values()))
assert isinstance(weights_cache_key, dict)
rep_cache_key = next(iter(weights_cache_key.values()))
buffer_view = target_buffer[buffer_id, 0, :, :lora_rank]
if rep is not None:
rep = rep * lora_adapter.scaling
rep = self._get_maybe_cached_weight_for_transfer(
pinned_layer_weights,
rep_cache_key,
rep,
)
load_lora_weight_tensor(buffer_view, rep)
# Zero beyond loaded rank — MoE kernel reads full max_rank
target_buffer[buffer_id, 0, :, lora_rank:].zero_()
else:
raise ValueError(
f"Unexpected weight format for shared outer down_proj_moe lora_B: "
f"type={type(weights)}, "
f"shape={weights.shape if isinstance(weights, torch.Tensor) else 'N/A'}"
)
# Zero beyond loaded rank — MoE kernel reads full max_rank.
target_buffer[buffer_id, 0, :, lora_rank:].zero_()
elif isinstance(weights, (torch.Tensor, dict)):
# Zero out slots this rank owns but the adapter
# doesn't fill (padded-out / out-of-rank experts);
# then scale+load the ones it does.
target_buffer[buffer_id].zero_()
for local_eid, w in self._iter_local_expert_weights(weights):
assert isinstance(weights_cache_key, (str, dict))
for (
local_eid,
w,
w_cache_key,
) in self._iter_local_expert_weights(
weights, weights_cache_key
):
if w is not None:
w = w * lora_adapter.scaling
w = self._get_maybe_cached_weight_for_transfer(
pinned_layer_weights,
w_cache_key,
w,
)
buffer_view = target_buffer[
buffer_id, local_eid, :, :lora_rank
]
load_lora_weight_tensor(buffer_view, w)
else:
buffer_view = target_buffer[buffer_id, :, :lora_rank]
if weights is not None:
assert isinstance(weights_cache_key, str)
weights = self._get_maybe_cached_weight_for_transfer(
pinned_layer_weights,
weights_cache_key,
weights,
)
load_lora_weight_tensor(buffer_view, weights)
if lora_adapter.embedding_layers:
org_vocab_size = self.base_hf_config.vocab_size
lora_added_tokens_size = lora_adapter.config.lora_added_tokens_size
pinned_embedding_layers = lora_adapter.pinned_embedding_layers
pinned_added_tokens_embeddings = lora_adapter.pinned_added_tokens_embeddings
# Only when LoRA is applied to the embedding layer will it have the extra-token issue that needs to be resolved.
# Load embeddings weights for extra tokens to buffer
if lora_adapter.added_tokens_embeddings:
@@ -1014,6 +1193,11 @@ class LoRAMemoryPool:
buffer_view = self.new_embeddings_buffer["input_embeddings"][
buffer_id, :lora_added_tokens_size
]
weights = self._get_maybe_cached_weight_for_transfer(
pinned_added_tokens_embeddings,
name,
weights,
)
load_lora_weight_tensor(buffer_view, weights)
# load vocab_emb and lm_head
@@ -1029,6 +1213,11 @@ class LoRAMemoryPool:
:lora_rank,
: (org_vocab_size + lora_added_tokens_size),
]
weights = self._get_maybe_cached_weight_for_transfer(
pinned_embedding_layers,
name,
weights,
)
load_lora_weight_tensor(buffer_view, weights)
elif (
target_module == "embed_tokens"
@@ -1042,6 +1231,11 @@ class LoRAMemoryPool:
buffer_view = self.embedding_B_buffer[target_module][
buffer_id, :, :lora_rank
]
lora_b_weights = self._get_maybe_cached_weight_for_transfer(
pinned_embedding_layers,
name,
lora_b_weights,
)
load_lora_weight_tensor(buffer_view, lora_b_weights)
elif (
@@ -1056,6 +1250,11 @@ class LoRAMemoryPool:
:lora_rank,
:,
]
weights = self._get_maybe_cached_weight_for_transfer(
pinned_embedding_layers,
name,
weights,
)
load_lora_weight_tensor(buffer_view, weights)
elif (
target_module == "lm_head"
@@ -1070,12 +1269,20 @@ class LoRAMemoryPool:
lora_b_weights = lora_lm_head_module.slice_lora_b_weights(
lora_b_weights, self.tp_rank
)
cache_key = append_cache_key_suffix(name, f"tp{self.tp_rank}")
else:
cache_key = name
buffer_view = self.lm_head_B_buffer[target_module][
buffer_id,
: lora_b_weights.shape[0],
:lora_rank,
]
lora_b_weights = self._get_maybe_cached_weight_for_transfer(
pinned_embedding_layers,
cache_key,
lora_b_weights,
)
load_lora_weight_tensor(buffer_view, lora_b_weights)
elif (
target_module == "lm_head"
+20
View File
@@ -83,6 +83,26 @@ class LoRAType(Enum):
LORA_B = 1
def copy_weight_into_buffer(
buffer_view: torch.Tensor,
weight: torch.Tensor,
) -> None:
"""
Copy a LoRA weight tensor into a destination buffer.
When a pinned CPU source has a dtype mismatch with a device destination,
cast on the destination device instead of doing the conversion on CPU.
"""
if weight.dtype == buffer_view.dtype:
buffer_view.copy_(weight, non_blocking=True)
return
if weight.device.type == "cpu" and buffer_view.device.type != "cpu":
weight = weight.to(device=buffer_view.device, non_blocking=True)
buffer_view.copy_(weight.to(dtype=buffer_view.dtype), non_blocking=True)
def get_hidden_dim(
module_name: str,
config: AutoConfig,
@@ -196,7 +196,11 @@ class TestIterLocalExpertWeightsDict(unittest.TestCase):
moe_use_local_expert_ids=False,
)
weights = {gid: torch.full((2,), float(gid)) for gid in range(4)}
got = {lid: w.tolist() for lid, w in pool._iter_local_expert_weights(weights)}
cache_keys = {gid: f"expert.{gid}" for gid in weights}
got = {
lid: w.tolist()
for lid, w, _ in pool._iter_local_expert_weights(weights, cache_keys)
}
self.assertEqual(
got, {0: [0.0, 0.0], 1: [1.0, 1.0], 2: [2.0, 2.0], 3: [3.0, 3.0]}
)
@@ -209,7 +213,11 @@ class TestIterLocalExpertWeightsDict(unittest.TestCase):
moe_use_local_expert_ids=True,
)
weights = {gid: torch.full((2,), float(gid)) for gid in range(8)}
got = {lid: w.tolist() for lid, w in pool._iter_local_expert_weights(weights)}
cache_keys = {gid: f"expert.{gid}" for gid in weights}
got = {
lid: w.tolist()
for lid, w, _ in pool._iter_local_expert_weights(weights, cache_keys)
}
# Rank 0 sees globals 0,1 remapped to locals 0,1.
self.assertEqual(got, {0: [0.0, 0.0], 1: [1.0, 1.0]})
@@ -221,7 +229,11 @@ class TestIterLocalExpertWeightsDict(unittest.TestCase):
moe_use_local_expert_ids=True,
)
weights = {gid: torch.full((2,), float(gid)) for gid in range(8)}
got = {lid: w.tolist() for lid, w in pool._iter_local_expert_weights(weights)}
cache_keys = {gid: f"expert.{gid}" for gid in weights}
got = {
lid: w.tolist()
for lid, w, _ in pool._iter_local_expert_weights(weights, cache_keys)
}
# Rank 3 sees globals 6,7 remapped to locals 0,1.
self.assertEqual(got, {0: [6.0, 6.0], 1: [7.0, 7.0]})
@@ -242,8 +254,12 @@ class TestIterLocalExpertWeightsDict(unittest.TestCase):
5: torch.full((2,), 5.0),
7: torch.full((2,), 7.0),
}
cache_keys = {gid: f"expert.{gid}" for gid in weights}
# Rank 2 owns globals 4, 5 -> locals 0, 1.
got = {lid: w.tolist() for lid, w in pool._iter_local_expert_weights(weights)}
got = {
lid: w.tolist()
for lid, w, _ in pool._iter_local_expert_weights(weights, cache_keys)
}
self.assertEqual(got, {0: [4.0, 4.0], 1: [5.0, 5.0]})
def test_no_experts_owned_yields_nothing(self):
@@ -258,7 +274,8 @@ class TestIterLocalExpertWeightsDict(unittest.TestCase):
)
# Only globals 4, 5 present (owned by rank 2).
weights = {4: torch.full((2,), 4.0), 5: torch.full((2,), 5.0)}
got = list(pool._iter_local_expert_weights(weights))
cache_keys = {gid: f"expert.{gid}" for gid in weights}
got = list(pool._iter_local_expert_weights(weights, cache_keys))
self.assertEqual(got, [])
@@ -275,9 +292,21 @@ class TestIterLocalExpertWeightsTensor(unittest.TestCase):
)
# [num_experts, rank, hidden] with values carrying the expert id.
weights = torch.arange(4 * 2 * 3, dtype=torch.float32).reshape(4, 2, 3)
got = [(lid, w.clone()) for lid, w in pool._iter_local_expert_weights(weights)]
self.assertEqual([lid for lid, _ in got], [0, 1, 2, 3])
for lid, w in got:
got = [
(lid, w.clone(), cache_key)
for lid, w, cache_key in pool._iter_local_expert_weights(weights, "weights")
]
self.assertEqual([lid for lid, _, _ in got], [0, 1, 2, 3])
self.assertEqual(
[cache_key for _, _, cache_key in got],
[
"weights#expert0",
"weights#expert1",
"weights#expert2",
"weights#expert3",
],
)
for lid, w, _ in got:
self.assertTrue(torch.equal(w, weights[lid]))
def test_rank1_of_ep2_sees_upper_half(self):
@@ -288,11 +317,18 @@ class TestIterLocalExpertWeightsTensor(unittest.TestCase):
moe_use_local_expert_ids=True,
)
weights = torch.arange(4 * 2 * 3, dtype=torch.float32).reshape(4, 2, 3)
got = [(lid, w.clone()) for lid, w in pool._iter_local_expert_weights(weights)]
got = [
(lid, w.clone(), cache_key)
for lid, w, cache_key in pool._iter_local_expert_weights(weights, "weights")
]
# Rank 1 of EP=2 with 4 experts owns globals 2, 3 -> locals 0, 1.
self.assertEqual([lid for lid, _ in got], [0, 1])
self.assertEqual([lid for lid, _, _ in got], [0, 1])
self.assertTrue(torch.equal(got[0][1], weights[2]))
self.assertTrue(torch.equal(got[1][1], weights[3]))
self.assertEqual(
[cache_key for _, _, cache_key in got],
["weights#expert2", "weights#expert3"],
)
def test_rank_with_partial_tensor_coverage(self):
"""Defensive: tensor has fewer experts than the expected local slice
@@ -310,7 +346,7 @@ class TestIterLocalExpertWeightsTensor(unittest.TestCase):
weights = torch.arange(6 * 2, dtype=torch.float32).reshape(6, 2)
# Note: this is 2D, not 3D -> should raise (sanity check).
with self.assertRaises(TypeError):
list(pool._iter_local_expert_weights(weights))
list(pool._iter_local_expert_weights(weights, "weights"))
class TestModuleLevelHelpers(unittest.TestCase):
@@ -701,6 +737,7 @@ class TestLoadBufferPassesMoeTpRankToSlice(unittest.TestCase):
torch.zeros(8, 4)
),
},
pinned_weights={},
)
]