sparkle/mlx_streaming/core/cache/resident_pool.py
fiser_jun 4745f264b2
2026-08-04 14:34:00 +08:00

636 lines
33 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""常驻专家池:每层一块连续 GPU 张量 + slot LRU支持按需增长、pin、GPU 侧重映射。
设计要点:
- 命中只返回槽位、不写池miss 只把单个专家原地写进它的槽位(`_write_slot`),避免拷贝整池。
- 池按需增长grow-on-demand起步 `_POOL_INIT_SLOTS` 行,工作集扩大时 ~1.5× 增至天花板
`cap_for(layer)` 才开始 LRU 淘汰。默认无 profile 即自动右尺寸,容量内永不超预算。
- `acquire_gpu`decode 热路径用 GPU 查找表做纯 GPU slot 重映射,命中层零 host 往返,
仅一次 miss 标志同步;真 miss 才回退 host 读盘路径。
"""
import os
from collections import OrderedDict, Counter
from typing import Dict, List
import mlx.core as mx
# 诊断门控staging/acquire_gpu 路径消费侧字节级真值校验(混合 ahead 损坏取证用,默认关)。
# 在 acquire_gpu 全命中快路径返回前,对本层真实路由命中的每个专家,把池槽字节与磁盘真值
# 逐 key 比对;不一致即「池槽装错字节」铁证。受 STG_VERIFY=1 控制,对主路径零影响(默认 off
_STG_VERIFY = os.environ.get("STG_VERIFY") == "1"
_stg_verify_state = {"ok": 0, "bad": 0, "printed": 0, "calls": 0, "first_bad_call": None}
from mlx_streaming import config
from mlx_streaming.core.cache import autopin
# Route 3 Phase 1 底座spec/dual 模式下池 buffer 改由 C++ 拥有(mx::allocator + no-op deleter)
# 地址进程内恒定、永不被 MLX donation/迁移,供侧区/demand 后台 pread 安全直写(替代消费侧 MLX scatter
# POOL_OWNED=0 可强制回退 mx.zeros(仅供 A/B 对照)。
_POOL_OWNED = os.environ.get("POOL_OWNED", "1") == "1"
# mx.Dtype -> C++ pool_owned_zeros 接受的 dtype 名。
_DTYPE_NAME = {
mx.uint32: "uint32", mx.uint16: "uint16", mx.uint8: "uint8",
mx.int32: "int32", mx.int16: "int16",
mx.bfloat16: "bfloat16", mx.float16: "float16", mx.float32: "float32",
}
def _owned_pool(sample: "Dict[str, mx.array]", n: int) -> "Dict[str, mx.array]":
"""用 C++-owned buffer 建 (n,*shape) 的 per-key 池数组(地址恒定,供 C++ 直写)。"""
import mlx_streaming.native_moe_ext as _N
out = {}
for k, v in sample.items():
name = _DTYPE_NAME.get(v.dtype)
if name is None:
raise RuntimeError(f"pool_owned_zeros 不支持 dtype {v.dtype} (key={k})")
out[k] = _N.pool_owned_zeros([int(n)] + [int(d) for d in v.shape], name)
return out
# 池按需增长(grow-on-demand)的初始物理行数。
# 起步小,工作集扩大时按 ~1.5× 增长,封顶 cap_for(默认=全局 capacity)。
# 好处:无需 profile 也能自动右尺寸(内存≈实际工作集),且对任何 prompt 自适应、容量内永不超预算。
_POOL_INIT_SLOTS = 16
class ResidentExpertPool:
"""每层一个连续常驻池:(capacity,*shape) 张量 + slot LRU。
命中只返回槽位、不写池miss 只把单个专家原地写进它的槽位_write_slot
loader(layer, e) -> Dict[str, mx.array],单个专家的参数(未堆叠)。
"""
def __init__(self, capacity: int, loader, layer_caps: "Dict[int, int] | None" = None,
spec_slots: int = 0, batch_loader=None, stacked_batch_loader=None,
spec_gens: int = 1):
self.capacity = capacity
self.loader = loader
# 可选批量加载器 batch_loader(layer, [ids]) -> {e: expert_dict}acquire 用它把本层所有
# miss 一次并行读8-worker pread取代逐专家串行 loader。None 时退回串行 loader。
self.batch_loader = batch_loader
# 可选「批量+预堆叠」加载器 stacked_batch_loader(layer, [ids]) -> {k:(N,*shape)}
# 在 batch_loader 基础上再把 6N 次碎 mx.array 物化折成每段一次,写池时直接整批 scatter
# 既省 frombuffer/mx.array 构造又省 _write_slots_batch 里的 mx.stack。优先级最高。
self.stacked_batch_loader = stacked_batch_loader
self.spec_slots = int(spec_slots) # >0侧区模式预分配 cap+spec、禁 grow
self.spec_gens = max(1, int(spec_gens)) # 侧区代数:双缓冲=2物理行=cap+spec_gens*spec_slots
# cap_for(layer) 是该层物理行数的「天花板」profile 指定则用之,否则=全局 capacity。
# 池按需增长(grow-on-demand):起步小、随工作集扩大增至天花板才开始 LRU 淘汰。
# 因此默认(无 profile)即自动右尺寸——内存≈实际工作集,对任意 prompt 自适应,
# 且天花板内永不超预算(超预算只是更慢,不崩)。增长偶发(预热期),稳态零拷贝。
self.layer_caps: "Dict[int, int]" = dict(layer_caps or {})
self._pools: "Dict[int, Dict[str, mx.array]]" = {}
self._slot_of: "Dict[int, OrderedDict[int, int]]" = {}
self._free: "Dict[int, list]" = {}
# 每层当前物理已分配行数(按需增长,≤ cap_for)。grow-on-demand 的核心状态。
self._alloc: "Dict[int, int]" = {}
# 每层 GPU 查找表(全局专家 id → slot,-1=不在池),供 acquire_gpu 在命中层
# 做纯 GPU 重映射(消除每层 .tolist 同步)。与 _slot_of 同步维护(仅在表已建时)。
self._slot_table: "Dict[int, mx.array]" = {}
# 每层 pinned 专家集合:预取进 resident pool 后不参与 LRU 驱逐。
# 用于 K>2 / 小槽位时保住热专家工作集,降低真实 miss。
self._pinned: "Dict[int, set[int]]" = {}
# 可选替换策略:默认 lfu短窗口频率 + LRU tie-break实测比 lru 命中更高EVICT_POLICY=lru 回退纯 LRU。
self.eviction_policy = config.evict_policy().lower()
self.lfu_decay_interval = config.lfu_decay_interval()
self._freq: "Dict[int, Counter[int]]" = {}
self._access_count: "Dict[int, int]" = {}
self.hits = 0
self.misses = 0
self.prefetch_hits = 0
self.prefetch_loads = 0
# GPU remap 路径取证:命中层走纯 GPU 快路径次数 vs 有 miss 回退 host 次数
self.gpu_fastpath = 0
self.gpu_fallback = 0
# 累计「分不到槽被迫落 0 号槽」的专家数(来自 demand_dual)。任何 >0 都意味着有前向
# 拿错专家权重算过 → 输出不可信,必须靠调小 PREFILL_CHUNK / 调大槽位消除。
self.unplaced = 0
# dual 路径真实区槽状态由 C++ demand_dual 唯一权威(无 opt-out)。spec 模式 + native 已编译即启用;
# 非 spec(spec_slots==0)或 native 缺失时保持 Python 权威路径(仅 prefill/host/非双源用)。
# 启用后 _slot_of/_free/_freq 在 dual 路径不再维护,resident_experts/_count 改查 C++ g_real(预取过滤要用)。
self._native_demand = False
if int(spec_slots) > 0:
try:
import mlx_streaming.native_moe_ext as _N
self._native_demand = hasattr(_N, "demand_dual")
except Exception:
self._native_demand = False
if self._native_demand:
# 复刻基线 8-worker 并行读demand miss 的 pread 派给 BgReader 并行执行(高优队列)。
import mlx_streaming.native_moe_ext as _N
try:
_N.bg_reader_start(int(os.environ.get("DEMAND_WORKERS", "8")), 0)
except Exception:
pass
if self._native_demand and os.environ.get("DEMAND_TIMING") == "1":
import atexit
import mlx_streaming.native_moe_ext as _N
_N.demand_timing_enable(True)
atexit.register(lambda: print(
"[DEMAND_TIMING ms] inds/pool/side_snap/real_lock/core/build =",
[round(x / 1e3, 1) for x in _N.demand_timings()], flush=True))
def cap_for(self, layer: int) -> int:
"""该层池容量profile 指定则用之(上限 capacity),否则用全局 capacity。"""
return min(self.layer_caps.get(layer, self.capacity), self.capacity)
def _ensure_layer(self, layer: int):
if layer not in self._slot_of:
self._slot_of[layer] = OrderedDict()
# 空闲槽列表懒填:首次 miss 分配池时再灌入,之后按需增长时追加。
# 注意:始终原地 mutate 这个 list,绝不重新绑定,确保 acquire 里持有的本地引用同步可见。
self._free[layer] = []
self._pinned[layer] = set()
self._freq[layer] = Counter()
self._access_count[layer] = 0
def _alloc_pool(self, layer: int, sample: Dict[str, mx.array]):
"""首次为某层分配池张量,起步 _POOL_INIT_SLOTS 行(不超过该层天花板)。"""
if self.spec_slots > 0:
n = self.cap_for(layer) + self.spec_gens * self.spec_slots # 预分配满(含 spec_gens 代侧区)
# spec/dual 模式:池由 C++ 拥有,侧区异步直写落此 bufferRoute 3 底座)。
if _POOL_OWNED:
self._pools[layer] = _owned_pool(sample, n)
else:
self._pools[layer] = {
k: mx.zeros((n,) + v.shape, dtype=v.dtype) for k, v in sample.items()}
else:
n = min(_POOL_INIT_SLOTS, self.cap_for(layer))
self._pools[layer] = {
k: mx.zeros((n,) + v.shape, dtype=v.dtype)
for k, v in sample.items()
}
self._alloc[layer] = n
real = self.cap_for(layer) if self.spec_slots > 0 else n
self._free[layer].extend(range(real)) # 原地 mutate,不重新绑定;侧区行不进 free永不被 LRU 分配/驱逐
def _grow_pool(self, layer: int, new_n: int):
"""把某层池物理行数扩到 new_n(封顶 cap_for),拼接保留已驻行的数据与 slot 索引。"""
if self.spec_slots > 0:
return # 侧区模式:预分配满,永不 grow保 C++ 写指针稳定)
old_n = self._alloc[layer]
new_n = min(new_n, self.cap_for(layer))
if new_n <= old_n:
return
old = self._pools[layer]
self._pools[layer] = {
k: mx.concatenate(
[v, mx.zeros((new_n - old_n,) + v.shape[1:], dtype=v.dtype)],
axis=0)
for k, v in old.items()
}
self._free[layer].extend(range(old_n, new_n)) # 新行追加为空闲,原地 mutate
self._alloc[layer] = new_n
def preallocate(self, layer: int, sample: "Dict[str, mx.array]", cap: int):
"""满 cap 预分配 typed 池并 eval 固定 data 指针;幂等。
"C++ 直写槽"机制使用:调用后该层不再 grow/不再走 _place_expert(MLX scatter)
槽位由调用方直接 mutate _slot_of + _set_table 管理、字节由 C++ pread 写入。
"""
self._ensure_layer(layer)
if layer in self._pools:
return
if _POOL_OWNED:
self._pools[layer] = _owned_pool(sample, cap) # C++ 拥有,地址恒定供直写
else:
self._pools[layer] = {
k: mx.zeros((cap,) + v.shape, dtype=v.dtype) for k, v in sample.items()
}
mx.eval(list(self._pools[layer].values())) # 固定 data 指针
self._alloc[layer] = cap
# 登记所有物理行为空闲,供共用分配器(_alloc_slot)按 free 优先分配;
# 已被 slot_of 占用的行不在此列(暖池复用时 slot_of 已满 → free 为空 → 走驱逐)。
occupied = set(self._slot_of[layer].values())
self._free[layer].extend(r for r in range(cap) if r not in occupied)
def allocated_slots(self, layer: int) -> int:
"""该层池张量当前物理行数(按需增长,≤ cap_for),用于内存核算/测试。"""
return self._alloc.get(layer, 0)
def _pool_owned(self, layer: int) -> bool:
"""该层池是否为 C++-owned bufferspec/dual 模式 + POOL_OWNED
owned 池禁止任何 MLX scatter会重分配 buffer、孤立侧区直写落池一律走 C++ 直写。"""
return _POOL_OWNED and self.spec_slots > 0
def _write_slot(self, layer: int, slot: int, expert: Dict[str, mx.array]):
if self._pool_owned(layer):
self._write_slots_batch(layer, [slot], [expert]) # 单槽也走 C++ 直写,避免 scatter
return
pool = self._pools[layer]
for k, v in expert.items():
pool[k][slot] = v # de-risk 选定:原地写,单槽、不拷贝整池
def _write_slots_batch(self, layer: int, slots: List[int],
experts: "List[Dict[str, mx.array]]") -> None:
"""把多个专家一次性写进各自槽位。
owned 池C++ memcpy 直写各行(无 MLX scatter保 buffer 地址恒定 → 侧区直写不被孤立)。
非 owned每个 key 一次 stack + fancy-index scatter原碎 kernel 最少化路径)。
"""
if not slots:
return
pool = self._pools[layer]
if self._pool_owned(layer):
import mlx_streaming.native_moe_ext as _N
keys = list(pool.keys())
pool_list = [pool[k] for k in keys]
srcs_flat = [e[k] for k in keys for e in experts] # key-major与 C++ pool_write_rows 约定一致
mx.eval(srcs_flat) # 物化源,固定 data 指针
_N.pool_write_rows(pool_list, srcs_flat, [int(s) for s in slots])
return
idx = mx.array(slots, dtype=mx.int32)
for k in pool:
pool[k][idx] = mx.stack([e[k] for e in experts], axis=0)
def resident_count(self, layer: int) -> int:
if self._native_demand: # 方案BC++ g_real 为真实区权威
import mlx_streaming.native_moe_ext as N
return int(N.real_region_count(int(layer)))
return len(self._slot_of.get(layer, ()))
def resident_experts(self, layer: int) -> set[int]:
"""返回该层 resident pool 里当前已有的专家 id 集合trace/probe + 侧区预取过滤用)。"""
if self._native_demand: # 方案B查 C++ g_real保预取过滤一致
import mlx_streaming.native_moe_ext as N
flat = N.real_region_contents(int(layer))
return {int(flat[i]) for i in range(0, len(flat), 2)}
return set(self._slot_of.get(layer, {}).keys())
def _bootstrap_dual_pool(self, layer: int) -> None:
"""方案B首次为该层预分配满(cap+spec_gens*spec)池、eval 固定 data 指针 + C++ real_init。幂等。
池结构(per-key typed 数组)由 Python 从一个样本专家的 shape 建;字节后续由 C++ demand pread 写入。
Python 侧 _slot_of/_free 在方案B dual 路径不再使用(真实区状态归 C++ g_real
"""
self._ensure_layer(layer)
if layer not in self._pools:
sample = self.loader(layer, 0) # 仅取 shape/dtype不写入池
cap = self.cap_for(layer)
n = cap + self.spec_gens * self.spec_slots
if _POOL_OWNED:
self._pools[layer] = _owned_pool(sample, n) # C++ 拥有,地址恒定供直写
else:
self._pools[layer] = {
k: mx.zeros((n,) + v.shape, dtype=v.dtype) for k, v in sample.items()
}
mx.eval(list(self._pools[layer].values())) # 固定 data 指针,供 C++ memcpy
self._alloc[layer] = n
import mlx_streaming.native_moe_ext as N
N.real_init(int(layer), int(cap))
def resident_lru_scores(self, layer: int) -> dict[int, float]:
"""返回 resident 专家的 LRU 新近度分数0=最久未用1=最近使用。"""
keys = list(self._slot_of.get(layer, {}).keys())
if not keys:
return {}
denom = max(1, len(keys) - 1)
return {int(e): i / denom for i, e in enumerate(keys)}
def _note_access(self, layer: int, expert_ids: List[int]) -> None:
if self.eviction_policy != "lfu":
return
self._ensure_layer(layer)
freq = self._freq[layer]
for e in expert_ids:
freq[int(e)] += 1
self._access_count[layer] += len(expert_ids)
if self.lfu_decay_interval > 0 and self._access_count[layer] >= self.lfu_decay_interval:
for e in list(freq):
freq[e] //= 2
if freq[e] <= 0:
del freq[e]
self._access_count[layer] = 0
def _choose_victim(self, layer: int, current: set[int]) -> int:
slot_of = self._slot_of[layer]
pinned = self._pinned.get(layer, set())
# 绝不驱逐当前请求(current)专家:否则 acquire 末尾按 slot_of 取槽会 KeyError。
# 只在「非 current 且非 pinned」中选受害者选不出来说明本次唯一专家+不可驱逐者已超容量,
# 由调用方(acquire 的 len(uniq)>cap 守卫)负责拦截,这里直接报清晰错误。
candidates = [e for e in slot_of if e not in pinned and e not in current]
if not candidates:
raise ValueError(
f"layer {layer} resident pool has no evictable non-current slot "
f"(capacity={self.cap_for(layer)}, pinned={len(pinned)}, current={len(current)})")
if self.eviction_policy != "lfu":
return candidates[0]
freq = self._freq.get(layer, Counter())
return min(enumerate(candidates), key=lambda x: (freq.get(x[1], 0), x[0]))[1]
def _alloc_slot(self, layer: int, e: int, current: "set[int]") -> "tuple[int, bool]":
"""为 e 分配/复用槽free 优先,否则驱逐非 pinned/非 current 的 LRU 最旧),
更新 slot_of + table但不写数据。返回 (slot, is_new)。
供 C++ 直写预取(prefetch_cpp)与 demand 写入(_place_expert)**共用**:保证两条路径
从同一套 _free/_slot_of/_slot_table 分配,绝不抢占同一物理槽。
"""
slot_of, free = self._slot_of[layer], self._free[layer]
if e in slot_of:
slot_of.move_to_end(e)
return slot_of[e], False
cap = self.cap_for(layer)
if not free and self._alloc.get(layer, 0) < cap:
cur = self._alloc[layer]
self._grow_pool(layer, cur + max(1, cur // 2))
if free:
slot = free.pop(0)
else:
evicted_e = self._choose_victim(layer, current)
slot = slot_of.pop(evicted_e)
self._clear_table(layer, evicted_e)
slot_of[e] = slot
slot_of.move_to_end(e)
self._set_table(layer, e, slot)
return slot, True
def _place_expert(self, layer: int, e: int, expert: Dict[str, mx.array],
current: "set[int] | None" = None) -> int:
"""把专家写进 resident pool必要时增长或驱逐非 pinned LRU返回 slot。"""
current = current or {e}
if layer not in self._pools:
self._alloc_pool(layer, expert)
slot, _ = self._alloc_slot(layer, e, current)
self._write_slot(layer, slot, expert)
return slot
def _place_experts(self, layer: int, ids: List[int],
experts: "Dict[int, Dict[str, mx.array]]",
current: "set[int]") -> List[int]:
"""批量把一组 miss 专家写进 resident pool共用槽分配器 + 单次批量 scatter。
与逐个 _place_expert 语义等价(同样的 grow/驱逐/槽分配顺序),
只把数据写入合并成 _write_slots_batch去掉每 token 上百次碎 scatter。
"""
if not ids:
return []
if layer not in self._pools:
self._alloc_pool(layer, experts[ids[0]])
slots = [self._alloc_slot(layer, e, current)[0] for e in ids]
self._write_slots_batch(layer, slots, [experts[e] for e in ids])
return slots
def _place_experts_stacked(self, layer: int, ids: List[int],
stacked: "Dict[str, mx.array]",
current: "set[int]") -> List[int]:
"""批量写一组 missstacked 为按 ids 顺序预堆叠的 {k:(N,*shape)}
每 key 直接一次 fancy-index scatter连 _write_slots_batch 的 mx.stack 都省了)。
"""
if not ids:
return []
if layer not in self._pools:
self._alloc_pool(layer, {k: v[0] for k, v in stacked.items()})
slots = [self._alloc_slot(layer, e, current)[0] for e in ids]
pool = self._pools[layer]
if self._pool_owned(layer):
import mlx_streaming.native_moe_ext as _N
keys = list(pool.keys())
pool_list = [pool[k] for k in keys]
stacked_list = [stacked[k] for k in keys]
mx.eval(stacked_list) # 物化预堆叠源
_N.pool_write_stacked(pool_list, stacked_list, [int(s) for s in slots])
return slots
idx = mx.array(slots, dtype=mx.int32)
for k in pool:
pool[k][idx] = stacked[k]
return slots
def prefetch_cpp(self, layer: int, expert_ids, submit_fn) -> list:
"""用 C++ 直写把预测专家投机预取进池(与 demand fallback 共用槽分配器)。
submit_fn(slot, expert):提交一次异步 C++ 后台读到该物理槽(调用方负责后续 wait
已驻 → 触摸 LRU 跳过;未驻 → 分配槽并 submit。返回新提交的 expert 列表。
投机性质:不计 hit/miss分得的槽可被后续 LRU 驱逐;无可驱逐槽时跳过该专家
(交给 demand fallback绝不抛错打断热路径。
"""
self._ensure_layer(layer)
if layer not in self._pools:
return [] # 池未建(首 token 预热)→ 全交 fallback
current = {int(e) for e in expert_ids}
slot_of = self._slot_of[layer]
submitted = []
for e in expert_ids:
e = int(e)
if e in slot_of:
slot_of.move_to_end(e)
continue
try:
slot, _ = self._alloc_slot(layer, e, current)
except ValueError:
continue # 无可驱逐槽 → 跳过,留给 demand fallback
submit_fn(slot, e)
submitted.append(e)
return submitted
def pin(self, layer: int, expert_ids: List[int]) -> None:
"""预取并钉住专家到 resident poolpinned 专家不被 LRU 驱逐。
pin 是显式预取,不计入 hit/miss 统计;后续 acquire/acquire_gpu 命中才计 hit。
"""
# dual 模式(demand_dual 唯一权威)已弃用 PIN_HOT:真实区归 C++ g_real,Python pinned 不再生效。
# 明确报错而非静默无效——请设 PIN_HOT=0。
if self._native_demand and list(expert_ids):
raise RuntimeError(
"PIN_HOT 在 demand_dual(dual 模式)下已弃用:真实区由 C++ g_real 权威、不支持 pin。"
"请设 PIN_HOT=0。")
loaded = {int(e): self.loader(layer, int(e)) for e in expert_ids}
self.pin_loaded(layer, loaded)
def pin_loaded(self, layer: int, experts: "Dict[int, Dict[str, mx.array]]") -> None:
"""把已加载专家钉进 resident pool避免 FileExpertStore.pin 重复读/持有两份。"""
uniq = list(dict.fromkeys(int(e) for e in experts))
cap = self.cap_for(layer)
self._ensure_layer(layer)
pinned = self._pinned[layer]
if len(pinned | set(uniq)) > cap:
raise ValueError(
f"pin 后 pinned 数量 {len(pinned | set(uniq))} > 该层池容量 {cap}")
slot_of = self._slot_of[layer]
for e in uniq:
if e in slot_of:
pinned.add(e)
slot_of.move_to_end(e)
continue
self._place_expert(layer, e, experts[e], current={e})
pinned.add(e)
def pin_loaded_dual(self, layer: int, experts: "Dict[int, Dict[str, mx.array]]") -> int:
"""dual(native demand)模式的 pinC++ g_real 是真实区唯一权威Python _slot_of 不生效),
必须 real_pin 注册pinned 免驱逐 + 从 free 头取真实区槽)再用 pool_write_rows 把字节
C++ 直写池行owned 池禁止 MLX scatter。返回成功 pin 的专家数。
仅在 _native_demand 下由 FileExpertStore.pin_batch 调用experts 为已加载的
{expert: {key: arr}}key 集/顺序与池一致——同源自 blob 段表)。
"""
import mlx_streaming.native_moe_ext as _N
ids = list(dict.fromkeys(int(e) for e in experts))
if not ids:
return 0
self._bootstrap_dual_pool(layer) # 建池 + real_init幂等
slots = _N.real_pin(int(layer), ids) # 登记 pinned + 分配真实区槽(-1=分不到)
ok = [(e, s) for e, s in zip(ids, slots) if s >= 0]
if not ok:
return 0
keys = list(self._pools[layer].keys())
pool_list = [self._pools[layer][k] for k in keys]
srcs_flat = [experts[e][k] for k in keys for e, _ in ok] # key-major同 _write_slots_batch
mx.eval(srcs_flat)
_N.pool_write_rows(pool_list, srcs_flat, [s for _, s in ok])
return len(ok)
def prefetch(self, layer: int, expert_ids: List[int]) -> None:
"""把专家预取进 resident pool但不计入 hit/miss且可被后续 LRU 驱逐。"""
uniq = list(dict.fromkeys(int(e) for e in expert_ids))
cap = self.cap_for(layer)
if len(uniq) > cap:
uniq = uniq[:cap]
self._ensure_layer(layer)
slot_of = self._slot_of[layer]
current = set(uniq)
for e in uniq:
if e in slot_of:
slot_of.move_to_end(e)
self.prefetch_hits += 1
continue
expert = self.loader(layer, e)
self._place_expert(layer, e, expert, current=current)
self.prefetch_loads += 1
def acquire(self, layer: int, expert_ids: List[int], protect: "set[int] | None" = None):
"""protect额外的「不可驱逐」专家集除本批 miss 外。dual 路径传入本前向全部 inds
确保落 miss 腾槽时不驱逐本前向仍需的命中专家(否则其 slot 被复用 → 消费侧 gather 错字节)。
仅影响驱逐候选,不改 hit/miss 记账口径。"""
# 唯一专家集合(保序去重):池只需同时容纳本次请求的唯一专家
uniq = list(dict.fromkeys(int(e) for e in expert_ids))
cap = self.cap_for(layer)
if len(uniq) > cap:
raise ValueError(
f"本次请求 {len(uniq)} 个唯一专家 > 该层池容量 {cap}")
self._ensure_layer(layer)
self._note_access(layer, uniq)
uniq_set = set(uniq)
# 驱逐保护集 = 本批 miss 调用方额外指定(如本前向全部 inds 命中专家)。
current = uniq_set if protect is None else (uniq_set | protect)
slot_of, free = self._slot_of[layer], self._free[layer]
# 先处理命中(触摸 LRU再把所有 miss 收成一批
misses = []
for e in uniq:
if e in slot_of:
self.hits += 1
slot_of.move_to_end(e) # 触摸为最近使用,避免本次内被驱逐
else:
misses.append(e)
if misses:
self.misses += len(misses)
if self.stacked_batch_loader is not None:
# 批量读 + 批量物化:每段一次 mx.array写池每 key 一次 scatter碎 kernel 最少)
stacked = self.stacked_batch_loader(layer, misses)
self._place_experts_stacked(layer, misses, stacked, current=current)
else:
if self.batch_loader is not None:
# 一次并行读本层所有 miss8-worker pread保序写池LRU/驱逐语义不变
loaded = self.batch_loader(layer, misses)
else:
loaded = {e: self.loader(layer, e) for e in misses}
# 批量写:每 key 仅一次 stacked scatter取代逐专家逐段碎 scatter
self._place_experts(layer, misses, loaded, current=current)
# slots 与原始 expert_ids 一一对应(含重复),便于直接 reshape 成 routing 索引
slots = [slot_of[int(e)] for e in expert_ids]
return self._pools[layer], slots
def _set_table(self, layer: int, e: int, slot: int):
t = self._slot_table.get(layer)
if t is not None:
t[int(e)] = slot # (num_experts,) 小张量原地改,代价可忽略
def _clear_table(self, layer: int, e: int):
t = self._slot_table.get(layer)
if t is not None:
t[int(e)] = -1
def _ensure_table(self, layer: int, num_experts: int) -> mx.array:
t = self._slot_table.get(layer)
if t is None or int(t.shape[0]) != num_experts:
self._ensure_layer(layer)
tab = [-1] * num_experts
for e, slot in self._slot_of[layer].items():
tab[int(e)] = int(slot)
t = mx.array(tab, dtype=mx.int32)
self._slot_table[layer] = t
return t
def acquire_gpu(self, layer: int, inds: mx.array, num_experts: int):
"""GPU 侧 slot 重映射:命中层零 host 往返(仅一次 miss 标志同步)。
返回 (pool_arrays, local)。inds 为路由结果(decode 时形如 (1,1,k))。
全命中 → local = 表[inds] 纯 GPU;有 miss → 回退既有 host acquire 读盘并维护表后重算 local。
"""
table = self._ensure_table(layer, num_experts)
local = mx.take(table, inds)
n_miss = int(mx.sum((local < 0).astype(mx.int32))) # 唯一一次 GPU→CPU 同步
if n_miss == 0:
self.gpu_fastpath += 1
self.hits += int(inds.size) # decode top-k 各专家互异 → size 即唯一命中数
if self.eviction_policy == "lfu": # 全命中层也计频:piggyback 上面 n_miss 的 drain,inds 已物化
flat = [int(i) for i in inds.reshape(-1).tolist()]
self._note_access(layer, flat)
autopin.note(layer, flat) # AUTOPIN 热度:搭同一物化点,零额外同步(AUTOPIN=0 立即返回)
return self._pools[layer], local
# 有 miss:回退 host 路径(读盘、写槽、维护表),再用更新后的表重算 local
self.gpu_fallback += 1
flat = [int(i) for i in inds.reshape(-1).tolist()]
autopin.note(layer, flat) # AUTOPIN 热度:miss 回退 flat 已物化,顺带计频
pool_arrays, _ = self.acquire(layer, flat)
local = mx.take(self._slot_table[layer], inds)
return pool_arrays, local
def verify_acquire_bytes(self, layer, inds, stg=None):
"""诊断(STG_VERIFY)acquire 后把本层真实路由命中专家的池槽字节与磁盘真值逐 key 比对。
发现不一致即打印 (call, layer, expert, slot, gen, key):池槽装错字节的铁证,
与 timing/gen 竞态对应。call 是全局校验调用序号(≈token 前向序,便于和首分歧 token 关联)。
stg可选 NativeStagingManager用于查该专家落池所用 gen。
"""
st = _stg_verify_state
st["calls"] += 1
call = st["calls"]
pool = self._pools.get(layer)
slot_of = self._slot_of.get(layer)
if pool is None or slot_of is None:
return
flat = {int(i) for i in inds.reshape(-1).tolist()}
for e in flat:
slot = slot_of.get(e)
if slot is None:
continue
try:
truth = self.loader(layer, e)
except Exception:
continue
bad_key = None
for k in pool:
if k not in truth:
continue
a = pool[k][slot]
b = truth[k]
if a.shape != b.shape or not bool(mx.all(a == b)):
bad_key = k
break
if bad_key is None:
st["ok"] += 1
else:
st["bad"] += 1
if st["first_bad_call"] is None:
st["first_bad_call"] = call
gen = None
if stg is not None:
gen = stg.placed_gen.get((layer, e))
if st["printed"] < 24:
st["printed"] += 1
print(f"[STG_VERIFY] BAD call={call} layer={layer} expert={e} "
f"slot={slot} gen={gen} key={bad_key} "
f"(ok={st['ok']} bad={st['bad']})", flush=True)
def hit_rate(self) -> float:
tot = self.hits + self.misses
return self.hits / tot if tot else 0.0