| |
| """ |
| Production-shape stability test based on real LLM configs. |
| 7 fixed worker processes, each bound to one NPU device. |
| """ |
| import numpy as np |
| import math |
| import subprocess |
| import os |
| import sys |
| import time |
| import multiprocessing as mp |
|
|
| MODELS = [ |
| ("Llama-3.1-8B / Qwen3-8B", 32, 128), |
| ("Llama-3-70B / Qwen3-32B", 64, 128), |
| ("Llama-3.1-405B", 128, 128), |
| ("Qwen3-235B-A22B", 64, 128), |
| ("DeepSeek-V3 K-side (MLA)",128, 192), |
| ("DeepSeek-V3 V-side (MLA)",128, 128), |
| ("Mistral / Gemma small", 32, 128), |
| ("Long-head experiment", 32, 64), |
| ("ULPs-aligned big head", 16, 256), |
| ] |
|
|
| Q_LENS = [1, 2, 4, 8] |
| KV_LENS_DECODE = [1, 64, 512, 1024, 2048, 4096, 8192, 16384] |
| KV_LENS_SPEC = [64, 1024, 2048, 4096, 8192] |
| KV_LENS_ODD = [1023, 1025, 1500, 3000] |
|
|
| DEVICES = [1, 2, 3, 4, 5, 6, 7] |
| ATOL = 0.001 |
| BUILD_DIR = os.path.abspath("./build") |
| IO_ROOT = "/tmp/smallq_test_io" |
| MAX_HALFS_PER_TEST = 200_000_000 |
|
|
|
|
| def build_test_cases(): |
| cases = [] |
| for model_name, nh, hd in MODELS: |
| for ql in Q_LENS: |
| if ql == 1: |
| kvs = KV_LENS_DECODE + KV_LENS_ODD |
| elif ql <= 4: |
| kvs = KV_LENS_SPEC + [1023, 3000] |
| else: |
| kvs = [64, 1024, 2048, 4096] |
| for kv in kvs: |
| total_halfs = nh * kv * hd + nh * ql * hd |
| if total_halfs > MAX_HALFS_PER_TEST: |
| continue |
| if ql >= 8 and hd >= 192: |
| continue |
| cases.append((model_name, nh, ql, hd, kv, "test_aclnn")) |
| return cases |
|
|
|
|
| def run_one(case_idx, device_id, io_dir, case): |
| (model_name, nh, ql, hd, kv, binary) = case |
| seed = case_idx % 65536 |
| np.random.seed(seed) |
| Q = np.random.randn(nh, ql, hd).astype(np.float16) |
| K = np.random.randn(nh, kv, hd).astype(np.float16) |
| V = np.random.randn(nh, kv, hd).astype(np.float16) |
| Q.tofile(f"{io_dir}/q.bin") |
| K.tofile(f"{io_dir}/k.bin") |
| V.tofile(f"{io_dir}/v.bin") |
|
|
| q32 = Q.astype(np.float32) |
| k32 = K.astype(np.float32) |
| v32 = V.astype(np.float32) |
| s = np.matmul(q32, k32.transpose(0, 2, 1)) / math.sqrt(hd) |
| s = s - s.max(axis=-1, keepdims=True) |
| p = np.exp(s) |
| p = p / p.sum(axis=-1, keepdims=True) |
| ref = np.matmul(p, v32).astype(np.float16).ravel() |
| del q32, k32, v32, s, p |
|
|
| env = os.environ.copy() |
| env.update({ |
| "NUM_HEADS": str(nh), "Q_LEN": str(ql), |
| "HEAD_DIM": str(hd), "KV_LEN": str(kv), |
| "DEVICE_ID": str(device_id), |
| "IO_DIR": io_dir, |
| }) |
| try: |
| r = subprocess.run([f"{BUILD_DIR}/{binary}"], env=env, |
| capture_output=True, text=True, timeout=180) |
| except subprocess.TimeoutExpired: |
| return (case_idx, device_id, model_name, nh, ql, hd, kv, binary, |
| "TIMEOUT", -1.0, "") |
| if r.returncode != 0: |
| err = (r.stderr or "")[-200:] |
| return (case_idx, device_id, model_name, nh, ql, hd, kv, binary, |
| "LAUNCH", -1.0, err) |
| out = np.fromfile(f"{io_dir}/output_o.bin", dtype=np.float16) |
| if out.size != ref.size: |
| return (case_idx, device_id, model_name, nh, ql, hd, kv, binary, |
| "SIZE", -1.0, f"got {out.size} vs ref {ref.size}") |
| diff = float(np.abs(out.astype(np.float32) - ref.astype(np.float32)).max()) |
| status = "OK" if diff <= ATOL else "DIFF" |
| return (case_idx, device_id, model_name, nh, ql, hd, kv, binary, |
| status, diff, "") |
|
|
|
|
| def worker_loop(device_id, task_q, result_q): |
| io_dir = f"{IO_ROOT}/dev{device_id}" |
| os.makedirs(io_dir, exist_ok=True) |
| while True: |
| item = task_q.get() |
| if item is None: |
| break |
| case_idx, case = item |
| result = run_one(case_idx, device_id, io_dir, case) |
| result_q.put(result) |
|
|
|
|
| def main(): |
| cases = build_test_cases() |
| total = len(cases) |
| print(f"=== Total test cases: {total} ===", flush=True) |
| print(f"=== Devices (one worker process each): {DEVICES} ===\n", flush=True) |
|
|
| task_q = mp.Queue() |
| result_q = mp.Queue() |
| for idx, case in enumerate(cases): |
| task_q.put((idx, case)) |
| for _ in DEVICES: |
| task_q.put(None) |
|
|
| workers = [] |
| for dev in DEVICES: |
| p = mp.Process(target=worker_loop, args=(dev, task_q, result_q)) |
| p.start() |
| workers.append(p) |
|
|
| pass_count = 0 |
| fail_count = 0 |
| failures = [] |
| t0 = time.time() |
|
|
| for done in range(total): |
| result = result_q.get() |
| (idx, dev, mn, nh, ql, hd, kv, bn, status, diff, err) = result |
| passed = (status == "OK") |
| if passed: |
| pass_count += 1 |
| else: |
| fail_count += 1 |
| failures.append(result) |
| flag = "✓" if passed else "✗" |
| suffix = f" {status}" + (f" err={err!r}" if err else "") |
| print(f"[{done+1:3d}/{total}] [dev{dev}] {flag} {bn:11s} {mn:30s} " |
| f"nh={nh:3d} ql={ql} hd={hd:4d} kv={kv:5d} " |
| f"diff={diff:.4f}{suffix}", flush=True) |
|
|
| for p in workers: |
| p.join() |
|
|
| dt = time.time() - t0 |
| print(f"\n=== DONE in {dt:.1f}s PASS={pass_count} FAIL={fail_count} ===") |
| if failures: |
| print("\nFailures:") |
| for f in failures: |
| print(f" {f}") |
| return 1 if fail_count > 0 else 0 |
|
|
|
|
| if __name__ == "__main__": |
| sys.exit(main()) |
|
|