Fix: init shared dict before using it, remove duplicate _output_buf
This commit is contained in:
@@ -119,8 +119,14 @@ class CuTeDSLMoERunner:
|
||||
for _ in range(self.num_experts)
|
||||
]
|
||||
|
||||
# Initialize shared buffers dict (if not already)
|
||||
device_key = str(self.device)
|
||||
if not hasattr(CuTeDSLMoERunner, '_shared_padded_bufs'):
|
||||
CuTeDSLMoERunner._shared_padded_bufs = {}
|
||||
if device_key not in CuTeDSLMoERunner._shared_padded_bufs:
|
||||
CuTeDSLMoERunner._shared_padded_bufs[device_key] = {}
|
||||
|
||||
# Padded x_sf buffers: SHARED across all runners (not per-layer)
|
||||
# Same reasoning as padded_hidden/activated — layers run sequentially.
|
||||
max_sf_rows = self.num_experts * self._max_chunks_per_expert * 128
|
||||
if 'xsf_l1' not in CuTeDSLMoERunner._shared_padded_bufs[device_key]:
|
||||
CuTeDSLMoERunner._shared_padded_bufs[device_key].update({
|
||||
@@ -142,33 +148,23 @@ class CuTeDSLMoERunner:
|
||||
self._l1_gsa_buf = torch.zeros(self.num_experts, dtype=torch.float32, device=self.device)
|
||||
self._l2_gsa_buf = torch.zeros(self.num_experts, dtype=torch.float32, device=self.device)
|
||||
|
||||
# Pre-allocated output buffer
|
||||
self._output_buf = torch.zeros(
|
||||
self.max_num_tokens, self.hidden_size, dtype=torch.bfloat16, device=self.device
|
||||
)
|
||||
|
||||
# Row indices for scale assembly (max_num_tokens * top_k slots)
|
||||
self._row_indices_buf = torch.arange(
|
||||
self.max_num_tokens * self.top_k, device=self.device
|
||||
)
|
||||
|
||||
# Padded hidden/activated: SHARED across all runners (not per-layer)
|
||||
# These are only used during run() which is sequential across layers.
|
||||
# Per-layer allocation would be 72 MB × 60 layers = 4.3 GB → OOM.
|
||||
max_rows_per_expert = self._max_chunks_per_expert * 128
|
||||
padded_max_slots = self.num_experts * max_rows_per_expert
|
||||
device_key = str(self.device)
|
||||
if not hasattr(CuTeDSLMoERunner, '_shared_padded_bufs'):
|
||||
CuTeDSLMoERunner._shared_padded_bufs = {}
|
||||
if device_key not in CuTeDSLMoERunner._shared_padded_bufs:
|
||||
CuTeDSLMoERunner._shared_padded_bufs[device_key] = {
|
||||
if 'hidden' not in CuTeDSLMoERunner._shared_padded_bufs[device_key]:
|
||||
CuTeDSLMoERunner._shared_padded_bufs[device_key].update({
|
||||
'hidden': torch.zeros(
|
||||
padded_max_slots, self.hidden_size, dtype=torch.bfloat16, device=self.device
|
||||
),
|
||||
'activated': torch.zeros(
|
||||
padded_max_slots, self.intermediate_size, dtype=torch.bfloat16, device=self.device
|
||||
),
|
||||
}
|
||||
})
|
||||
self._shared_bufs = CuTeDSLMoERunner._shared_padded_bufs[device_key]
|
||||
|
||||
# Padded expert offsets buffer: [0, max_rows, 2*max_rows, ...] (fixed)
|
||||
|
||||
Reference in New Issue
Block a user