""" P800 Manual Timing Module for SGLang Monkey-patches DeepSeek-V4 decoder layers to record per-op timing. Bypasses broken PyTorch Profiler/CUPTI on P800 XPU. Usage: set SGLANG_TIMING_ENABLED=1 in container env """ import os import time import json import atexit import signal import functools import threading _lock = threading.Lock() _timing = {} _write_counter = 0 _WRITE_INTERVAL = 5000 # write every N records _phase = threading.local() # per-thread: "prefill" or "decode" def _get_output_path(): base = os.environ.get("SGLANG_TIMING_OUTPUT", "/tmp/p800_timing_results.json") pid = os.getpid() name, ext = os.path.splitext(base) return f"{name}_{pid}{ext}" def _write_to_file(path, summary, raw_count): result = { "summary": summary, "config": { "timestamp": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), "pid": os.getpid(), }, "raw_count": raw_count, } try: os.makedirs(os.path.dirname(path) or ".", exist_ok=True) with open(path, "w") as f: json.dump(result, f, indent=2) except Exception as exc: print(f"[P800_TIMING] Write error: {exc}", flush=True) def _compute_summary(): summary = {} raw_count = {} for key, samples in _timing.items(): if not samples or not isinstance(samples, list): continue vals = [s.get("t", 0) for s in samples if isinstance(s.get("t"), (int, float))] if not vals: continue vals.sort() n = len(vals) raw_count[key] = n summary[key] = { "count": n, "total_ms": round(sum(vals), 3), "avg_ms": round(sum(vals) / n, 4), "p50_ms": round(vals[n // 2], 4), "p95_ms": round(vals[int(n * 0.95)], 4), "p99_ms": round(vals[int(n * 0.99)], 4), } return summary, raw_count def _flush(): with _lock: summary, raw_count = _compute_summary() if summary: _write_to_file(_get_output_path(), summary, raw_count) def _record(key, t_ms, extra=None, phase=None): global _write_counter if phase is None: phase = getattr(_phase, "mode", "unknown") full_key = f"{key}_{phase}" if phase != "unknown" else key with _lock: e = {"t": round(t_ms, 4), "at": time.time(), "phase": phase} if extra: e.update(extra) _timing.setdefault(full_key, []).append(e) _write_counter += 1 if _write_counter % _WRITE_INTERVAL == 0: _flush() def _sync(): try: import torch torch.cuda.synchronize() except Exception: pass @atexit.register def _write(): _flush() def _sigterm_handler(signum, frame): _flush() signal.signal(signal.SIGTERM, signal.SIG_DFL) os.kill(os.getpid(), signal.SIGTERM) try: signal.signal(signal.SIGTERM, _sigterm_handler) except Exception: pass def _make_timed_fn(orig_fn, key): """Create a timed wrapper for a given function. key can be a string or a callable(args, kwargs) -> string.""" @functools.wraps(orig_fn) def timed_fn(*args, **kwargs): _sync() t0 = time.time() try: return orig_fn(*args, **kwargs) finally: _sync() k = key(*args, **kwargs) if callable(key) else key _record(k, (time.time() - t0) * 1000) return timed_fn def patch(): """Apply monkey-patches to DeepSeek-V4 model classes.""" try: from sglang.srt.models import deepseek_v4 except ImportError: return False # ── Layer forward (split prefill vs decode) ── orig_layer_fwd = deepseek_v4.DeepseekV4DecoderLayer.forward @functools.wraps(orig_layer_fwd) def timed_layer_fwd(self, *args, **kwargs): # Determine prefill vs decode from forward_batch or hidden_states shape fb = kwargs.get("forward_batch", args[3] if len(args) > 3 else None) hs = kwargs.get("hidden_states", args[1] if len(args) > 1 else None) # Determine prefill vs decode from forward_batch flags # is_prefill_only: true if ALL reqs in batch are prefill # is_extend_in_batch: true if any req in batch has extend tokens # Fallback: check extend_num_tokens > 0 is_prefill = False if fb is not None: if getattr(fb, "is_prefill_only", False): is_prefill = True elif getattr(fb, "is_extend_in_batch", False): is_prefill = True elif getattr(fb, "extend_num_tokens", None) is not None and fb.extend_num_tokens > 0: is_prefill = True _phase.mode = "prefill" if is_prefill else "decode" _sync() t0 = time.time() try: return orig_layer_fwd(self, *args, **kwargs) finally: _sync() phase = _phase.mode dt = (time.time() - t0) * 1000 _record("layer_ms", dt, phase=phase, extra={"layer": self.layer_id}) deepseek_v4.DeepseekV4DecoderLayer.forward = timed_layer_fwd # ── HC Pre / Post ── for hc_name in ["hc_pre", "hc_post"]: fn = getattr(deepseek_v4.DeepseekV4DecoderLayer, hc_name, None) if fn is not None: setattr(deepseek_v4.DeepseekV4DecoderLayer, hc_name, _make_timed_fn(fn, f"{hc_name}_ms")) print(f"[P800_TIMING] Patched {hc_name}", flush=True) # ── Attention forward (MQALayer) ── for name in ["MQALayer", "RadixAttention", "DeepseekV4Attention", "DeepseekV4FlashAttention"]: cls = getattr(deepseek_v4, name, None) if cls is not None and hasattr(cls, "forward"): cls.forward = _make_timed_fn(cls.forward, "attention_ms") print(f"[P800_TIMING] Patched attention: {name}", flush=True) break # ── MQALayer._forward_prepare (QKV projection) ── cls = getattr(deepseek_v4, "MQALayer", None) if cls is not None and hasattr(cls, "_forward_prepare"): cls._forward_prepare = _make_timed_fn(cls._forward_prepare, "qkv_prepare_ms") print(f"[P800_TIMING] Patched MQALayer._forward_prepare", flush=True) # ── RMSNorm (both input_layernorm and post_attention_layernorm) ── for rms_name in ["RMSNorm", "DeepseekRefRMSNorm"]: cls = getattr(deepseek_v4, rms_name, None) if cls is not None and hasattr(cls, "forward"): cls.forward = _make_timed_fn(cls.forward, "rmsnorm_ms") print(f"[P800_TIMING] Patched RMSNorm: {rms_name}", flush=True) break # ── MLP/MoE (deepseek_v2.DeepseekV2MoE) ── try: from sglang.srt.models import deepseek_v2 cls = deepseek_v2.DeepseekV2MoE if hasattr(cls, "forward"): cls.forward = _make_timed_fn(cls.forward, "mlp_ms") print(f"[P800_TIMING] Patched MoE: DeepseekV2MoE", flush=True) moe_patched = True # ── MoE sub-operators ── # Gate if hasattr(cls, "gate") and hasattr(cls.gate, "forward"): cls.gate.forward = _make_timed_fn(cls.gate.forward, "moe_gate_ms") print(f"[P800_TIMING] Patched MoE gate", flush=True) # Shared experts if hasattr(cls, "_forward_shared_experts"): orig_shared = cls._forward_shared_experts cls._forward_shared_experts = _make_timed_fn(orig_shared, "moe_shared_ms") print(f"[P800_TIMING] Patched MoE shared_experts", flush=True) # forward_normal (routed experts kernel) if hasattr(cls, "forward_normal"): # The actual name might be klx_forward_normal for fn_name in ["forward_normal", "klx_forward_normal"]: fn = getattr(cls, fn_name, None) if fn is not None: setattr(cls, fn_name, _make_timed_fn(fn, "moe_routed_ms")) print(f"[P800_TIMING] Patched MoE {fn_name}", flush=True) break elif hasattr(cls, "forward_mega_moe"): cls.forward_mega_moe = _make_timed_fn(cls.forward_mega_moe, "moe_routed_ms") print(f"[P800_TIMING] Patched MoE forward_mega_moe", flush=True) except Exception as e: print(f"[P800_TIMING] MoE patch error: {e}", flush=True) moe_patched = False if not moe_patched: for name in ["FusedMoE", "DeepseekV4MoE", "MoE"]: cls = getattr(deepseek_v4, name, None) if cls is not None and hasattr(cls, "forward"): cls.forward = _make_timed_fn(cls.forward, "mlp_ms") print(f"[P800_TIMING] Patched MoE: {name} (fallback)", flush=True) break # ── All-reduce ── try: from sglang.srt.layers.communicator import Communicator if hasattr(Communicator, "all_reduce"): Communicator.all_reduce = _make_timed_fn(Communicator.all_reduce, "all_reduce_ms") print(f"[P800_TIMING] Patched all_reduce", flush=True) except Exception as e: print(f"[P800_TIMING] all_reduce patch failed: {e}", flush=True) print("[P800_TIMING] Patches applied: layer, hc_pre, hc_post, attention, " "qkv_prepare, rmsnorm, moe, gate, shared, routed", flush=True) return True def _wait_and_patch(): """Poll for deepseek_v4 module and apply patches.""" import sys import threading as _thr enabled = os.environ.get("SGLANG_TIMING_ENABLED", "0") if enabled != "1": return def _poll(): while True: if 'sglang.srt.models.deepseek_v4' in sys.modules: patch() return _thr.Event().wait(0.5) _thr.Thread(target=_poll, daemon=True).start() _wait_and_patch()