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

238 lines
14 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.

"""VirtualPool预取调度器 + 双源双缓冲协调器(统一收口)。
一个对象承担两职,因为 block.py 用同一个 `self._vpool` 属性:
1. ahead 调度(两种模式都用):`ahead_for` / `target_for` 决定「第 L 层应预读哪层」——
早层用小 ahead 保召回、cutoff 起用大 ahead 抢时序。`_native_fused_prefetch` 靠它选目标层。
2. 双源双缓冲协调(仅 ZEROCOPY_DUAL_SOURCE 模式):`begin_forward` / `read_gen` /
`fill_gen` / `acquire` / `prefetch`。侧区分两「代(gen)」物理行每前向读上一代、fill 写
另一代,根除「本前向消费 gather 读的物理行在本前向 eval 期被 fill 覆盖」的竞态
(见 spec 2026-06-25-qwen-virtual-pool-double-buffer。对外只暴露「专家→物理行」单次
gather 接口,消费者每层零 host 同步。
构造两种签名并存(互斥使用):
- 调度器VirtualPool(num_layers=.., cutoff=.., ahead_lo=.., ahead_hi=..)
- 协调器VirtualPool(resident, staging, spec_slots)dual-source 下再补调度参数即可两职合一。
"""
import mlx.core as mx
from mlx_streaming.core.cache import autopin
# 方案B STG_VERIFY 校验累计态(诊断用,默认路径不触及)。
_stg_verify_state = {"ok": 0, "bad": 0, "printed": 0, "calls": 0}
class VirtualPool:
def __init__(self, resident=None, staging=None, spec_slots=None, *,
num_layers=None, cutoff=None, ahead_lo=None, ahead_hi=None, store=None):
# --- 双源协调resident/staging 存在时启用)---
self._rp = resident
self._stg = staging
self._store = store # acquire_host 超容量 fetch 回退用(blob 感知);调度器形态可为 None
self._spec = int(spec_slots) if spec_slots is not None else 0
self._gen = 0
self._last_layer = -1 # 前向边界检测:层号回绕(<= 上次) 即新前向
# --- ahead 调度 ---
self._num_layers = int(num_layers) if num_layers is not None else 0
self._cutoff = int(cutoff) if cutoff is not None else 0
self._a_lo = max(1, int(ahead_lo)) if ahead_lo is not None else 1
self._a_hi = max(1, int(ahead_hi)) if ahead_hi is not None else 1
# ---- ahead 调度 ----
def ahead_for(self, src_layer: int) -> int:
# cutoff早层用 lo保召回cutoff 起用 hi保时序
return self._a_lo if int(src_layer) < self._cutoff else self._a_hi
def target_for(self, src_layer: int) -> int:
# 目标层 = src+ahead。末层无可预读 → 返回 0跳过
# 不再 clamp 到末层clamp 会让多个源层同时预读同一末层,对该层每前向 submit 次数 >
# staging ring其环形 buffer 在惰性切片被本前向 eval 消费前就被后续 submit 完成回调覆盖成
# 别的专家字节 → 池槽装错字节。越界即跳过(交给 demand 读真值),保证每层每前向至多 1 次 submit。
L = int(src_layer)
if L >= self._num_layers - 1:
return 0
tgt = L + self.ahead_for(L)
if tgt > self._num_layers - 1: # 越界(本会触发 clamp 堆叠)→ 跳过
return 0
return tgt
# ---- 双源双缓冲协调 ----
def begin_forward(self, layer_idx: int):
"""每个 MoE 块 __call__ 开头调;层号回绕(<= 上次) 判为新前向 → 代 +1。
稳健:不依赖首个 MoE 层是 layer 0、也不要求 MoE 层连续。"""
if layer_idx <= self._last_layer:
self._gen += 1
# AUTOPIN:dual decode 无 Python 计频点,用前向边界驱动周期落盘(AUTOPIN=0 立即返回)。
autopin.tick()
# 新前向开头排空上一前向提交的侧区 fillC++-owned 池 buffer 由后台异步直写,
# 若消费前 fill 未写完GPU gather 会读到半写侧区行DUAL_VERIFY BAD
# drain 阻塞到在途 fill 全部写完 → 本前向要消费的侧区行字节必已就绪。
if self._spec > 0:
import mlx_streaming.native_moe_ext as _N
_N.sideregion_drain()
self._last_layer = layer_idx
def _gens(self) -> int:
# 代数取自常驻池;单代(=1)时读=填=0(持久 LFU 单区),双代(=2)时交替(%2==&1)。
g = getattr(self._rp, "spec_gens", 2) if self._rp is not None else 2
return max(1, int(g))
def read_gen(self) -> int:
return (self._gen - 1) % self._gens() # 读上一前向填好的代;单代恒 0
def fill_gen(self) -> int:
return self._gen % self._gens() # fill 写本代;单代恒 0
def acquire(self, layer, inds, num_experts, *, seq_len=None, layer_cap=None):
"""统一取用入口GPU-remap 路径):对外呈现「所有专家都在」的视角。
返回 (pool_arrays, local, n_experts),计算侧零分支:
- dual有侧区 staging 且 spec>0C++ demand_dual 唯一权威(真实区 侧区读代单次 gather
n_experts = layer_cap + spec_gens*spec_slots。
- 非 dual GPU-remapacquire_gpun_experts = layer_cap。
host/fetch 路径见 acquire_host其输入是 host 侧已 .tolist 的 flat语义不同故分开。
"""
cap = int(layer_cap) if layer_cap is not None else self._rp.cap_for(layer)
if self._stg is not None and self._spec > 0:
side_gen = self.read_gen()
# C++ demand_dual 是双源 decode 真实区的唯一权威(每层 1 次 inds 同步、零主线程落池/记账)。
# 无 Python 退路native 未编译 demand_dual → 明确报错decode 依赖 native 全套能力)。
if not getattr(self._rp, "_native_demand", False):
raise RuntimeError(
"双源 decode 需要 native demand_dual已成唯一权威但未检出。请编译 native_moe_ext。")
pool, local = self._acquire_native(layer, inds, side_gen, cap)
n_exp = cap + self._rp.spec_gens * self._rp.spec_slots
return pool, local, n_exp
pool, local = self._rp.acquire_gpu(layer, inds, num_experts)
return pool, local, cap
def _native_meta(self, layer):
"""缓存每层 demand_dual 的不变入参pool_list/seg_nbytes/path避免逐层重建 host 胶水。"""
cache = getattr(self, "_nd_meta", None)
if cache is None:
cache = self._nd_meta = {}
m = cache.get(layer)
if m is None:
stg = self._stg
segs = stg.src._segs # (proj, tensor, dt, shape, nb),与池 key 同序
pool_list = [self._rp._pools[layer][f"{p}.{t}"] for p, t, *_ in segs]
m = (pool_list, [int(nb) for *_, nb in segs],
f"{stg.src.dir}/layer{int(layer):02d}.blob", int(stg.stride))
cache[layer] = m
return m
def _acquire_native(self, layer, inds, side_gen, cap):
"""方案B 取用:委派 C++ demand_dual每层 1 次 inds 同步 + 并行 worker pread 落池),
更新 rp 统计计数(供报告口径一致)。"""
import mlx_streaming.native_moe_ext as N
rp = self._rp
rp._bootstrap_dual_pool(layer) # 首次建池 + real_init幂等
# 方案B 容量前提miss 只能落真实区的 cap 槽,其中 pinned 永不可驱逐 → 可安放上限 cappinned。
# 超了 C++ 会把专家落 0 号槽(拿别的专家权重算)→ 逐位不正确。调用方block.py 的分流判据)
# 负责在超容量前把该前向送去 host/fetch这里只做一次性告警兜底。
_placeable = int(cap) - int(N.real_pinned_count(int(layer)))
if int(inds.size) > _placeable and not getattr(self, "_overcap_warned", False):
self._overcap_warned = True
import sys
print(f"[DEMAND_DUAL] 警告inds.size={int(inds.size)} > 可安放槽 {_placeable}"
f"(cap={int(cap)} pinned={int(cap) - _placeable}),真实区超容量,逐位将不正确。"
f"请调小 PREFILL_CHUNK或调大 EXPERT_SLOTS / 调小 AUTOPIN_BUDGET_FRAC。",
file=sys.stderr, flush=True)
pool_list, seg_nbytes, path, stride = self._native_meta(layer)
local = N.demand_dual(inds, pool_list, seg_nbytes, int(layer), int(side_gen), path,
stride, int(cap),
rp.eviction_policy == "lfu", int(rp.lfu_decay_interval))
st = N.demand_last_stats() # [hitpos, misspos, loads, fallback01, unplaced]
rp.hits += st[0]
rp.misses += st[2]
if len(st) > 4:
rp.unplaced += st[4] # >0 即有前向落了 0 号槽,输出不可信(/api/stats 暴露)
if st[3] == 0:
rp.gpu_fastpath += 1
else:
rp.gpu_fallback += 1
from mlx_streaming import config
if config.stg_verify(): # 诊断:方案B 池字节逐 key 真值校验(默认关)
self._verify_native_bytes(layer, inds, local)
return rp._pools[layer], local
def _verify_native_bytes(self, layer, inds, local):
"""诊断(STG_VERIFY方案B):校验「字节落池不变量」——真实区每个占用槽的池字节 == 该槽当前
C++ 属主专家(g_real)的 blob 真值。这是 C++ 接管落池的字节等价铁证;发现不一致即池装错字节。
注:不以 local→expert 为判据local 可能因跨调用/多模型共享 g_real 而滞后于 g_real属路由级
问题、非落池字节问题;逐位权威信号是 e2e n_mismatch
"""
st = _stg_verify_state
st["calls"] += 1
pool = self._rp._pools.get(layer)
if pool is None:
return
import mlx_streaming.native_moe_ext as N
stg = self._stg
path = f"{stg.src.dir}/layer{int(layer):02d}.blob"
segs = stg.src._segs
flat = N.real_region_contents(int(layer)) # [expert0,slot0,expert1,slot1,...]
for j in range(0, len(flat), 2):
e, slot = flat[j], flat[j + 1]
raw = N.blob_load(path, mx.array([e], dtype=mx.uint32), int(stg.stride))[0]
bad, off = None, 0
for p, t, dt, shape, nb in segs:
k = f"{p}.{t}"
pv = pool[k][slot].reshape(-1).view(mx.uint8)
if not bool(mx.all(pv == raw[off:off + nb])):
bad = k
break
off += nb
if bad is None:
st["ok"] += 1
else:
st["bad"] += 1
if st["printed"] < 12:
st["printed"] += 1
print(f"[STG_VERIFY-DUAL] BAD 落池字节错 call={st['calls']} layer={layer} "
f"expert={e} slot={slot} key={bad} (ok={st['ok']} bad={st['bad']})", flush=True)
def acquire_host(self, layer, flat, inds_shape, inds_dtype, layer_cap):
"""host/fetch 路径收口prefill/大 seq 或关 GPU-remapflat 为 host 侧路由 id 列表。
返回 (pool_arrays, local, n_experts),与 block.py 原 host/fetch 分支逐元素等价:
- uniq <= capacquire(flat)local 为槽位n_experts = layer_cap。
- uniq > capfetch(uniq_sorted)local 为 remap 到 [0,uniq) 连续索引n_experts = uniq 数。
native demand(dual)下真实区槽状态归 C++ g_realPython 的 _slot_of/_free 从不填充,
走 self._rp.acquire 必然在 _choose_victim 抛 "no evictable non-current slot"
故 dual 下一律走 fetch 旁路(不进池、正确但慢)——这条路径本就是超容量前向的兜底。
"""
import mlx.core as mx
cap = int(layer_cap)
uniq_set = set(flat)
# 入池的充分条件是 uniq ≤ cap pinned:acquire 把本次全部唯一专家列为不可驱逐(current),
# 腾槽时只能挑「非 pinned 且非 current」的受害者,少一个都会在 _choose_victim 抛
# "no evictable non-current slot"。pinned 是那层实际钉死数,不是拍出来的固定余量——
# 写死余量会在 cap 小于余量时把所有请求推进 fetch 慢路径(每次重读盘)。
_pinned_n = len(getattr(self._rp, "_pinned", {}).get(layer, ()) or ())
if (not getattr(self._rp, "_native_demand", False)
and len(uniq_set) <= max(1, cap - _pinned_n)):
pool, slots = self._rp.acquire(layer, flat)
local = mx.array(slots, dtype=inds_dtype).reshape(inds_shape)
return pool, local, cap
uniq_sorted = sorted(uniq_set)
remap = {g: i for i, g in enumerate(uniq_sorted)}
local = mx.array([remap[i] for i in flat], dtype=inds_dtype).reshape(inds_shape)
# 超容量(或 dual:真实区槽状态归 C++ g_real,Python _slot_of 从不填充 → acquire 必抛)
# 走 fetch 旁路:不进池、按 uniq 现堆一份参与计算,正确但慢。这条路径本就是兜底。
# fetch 由 blob 感知的 FileExpertStore 提供;ResidentExpertPool 本身没有,故优先用注入的 store。
src = self._store if self._store is not None else self._rp
if not hasattr(src, "fetch"):
raise RuntimeError("VirtualPool.acquire_host 超容量回退需要带 fetch 的 store(构造时传入)")
fetched = src.fetch(layer, uniq_sorted)
return fetched, local, len(uniq_sorted)
def prefetch(self, layer, pred, resident, pool_list):
"""向 fill 代 submit 预读base_row = cap_for(layer) + fill_gen*spec。"""
g = self.fill_gen()
base = self._rp.cap_for(layer) + g * self._spec
return self._stg.submit_pool_sideregion(layer, pred, resident, pool_list, base, gen=g)