|
| 1 | +import argparse |
| 2 | +import subprocess |
| 3 | +import time |
| 4 | +import urllib.request |
| 5 | +import urllib.error |
| 6 | +import json |
| 7 | +import re |
| 8 | +import sys |
| 9 | +import os |
| 10 | + |
| 11 | +CONFIGS = [ |
| 12 | + {"name": "Dense/Vanilla", "flags": []}, |
| 13 | + {"name": "SSD Stream", "flags": ["--stream-experts"]}, |
| 14 | + {"name": "TurboQuant", "flags": ["--turbo-kv"]}, |
| 15 | + {"name": "SSD + TurboQuant", "flags": ["--stream-experts", "--turbo-kv"]} |
| 16 | +] |
| 17 | + |
| 18 | +SWIFTLM_PATH = ".build/arm64-apple-macosx/release/SwiftLM" |
| 19 | + |
| 20 | +def poll_health(port=5413, timeout=30): |
| 21 | + start = time.time() |
| 22 | + url = f"http://127.0.0.1:{port}/health" |
| 23 | + while time.time() - start < timeout: |
| 24 | + try: |
| 25 | + r = urllib.request.urlopen(url) |
| 26 | + if r.getcode() == 200: |
| 27 | + return True |
| 28 | + except: |
| 29 | + pass |
| 30 | + time.sleep(1) |
| 31 | + return False |
| 32 | + |
| 33 | +def make_request_stream(prompt_len, max_tokens, port=5413): |
| 34 | + prompt = "apple " * int(prompt_len * 0.75) |
| 35 | + data = json.dumps({ |
| 36 | + "messages": [{"role": "user", "content": prompt}], |
| 37 | + "max_tokens": max_tokens, |
| 38 | + "temperature": 0.0, |
| 39 | + "stream": True |
| 40 | + }).encode('utf-8') |
| 41 | + |
| 42 | + req = urllib.request.Request( |
| 43 | + f"http://127.0.0.1:{port}/v1/chat/completions", |
| 44 | + data=data, |
| 45 | + headers={'Content-Type': 'application/json'} |
| 46 | + ) |
| 47 | + |
| 48 | + ttft = None |
| 49 | + start = time.time() |
| 50 | + tokens = 0 |
| 51 | + try: |
| 52 | + with urllib.request.urlopen(req, timeout=120) as response: |
| 53 | + for line in response: |
| 54 | + line = line.decode('utf-8').strip() |
| 55 | + if line.startswith("data: ") and line != "data: [DONE]": |
| 56 | + if ttft is None: |
| 57 | + ttft = time.time() - start |
| 58 | + tokens += 1 |
| 59 | + total_time = time.time() - start |
| 60 | + gen_time = total_time - ttft if ttft else 0 |
| 61 | + tps = (tokens - 1) / gen_time if gen_time > 0 and tokens > 1 else 0 |
| 62 | + return True, ttft, tps |
| 63 | + except Exception as e: |
| 64 | + print(f"Request failed: {e}") |
| 65 | + return False, 0, 0 |
| 66 | + |
| 67 | +def extract_base_memory(log_path): |
| 68 | + try: |
| 69 | + with open(log_path, 'r') as f: |
| 70 | + for line in f: |
| 71 | + if "Memory strategy: FULL GPU" in line: |
| 72 | + m = re.search(r"\(([0-9.]+)GB model", line) |
| 73 | + if m: return f"{m.group(1)} GB" |
| 74 | + except: pass |
| 75 | + return "N/A" |
| 76 | + |
| 77 | +def extract_real_memory(log_path): |
| 78 | + try: |
| 79 | + with open(log_path, 'r') as f: |
| 80 | + log_data = f.read() |
| 81 | + m = re.findall(r"OS_RAM=([0-9.]+)", log_data) |
| 82 | + if m: return f"{m[-1]} GB" |
| 83 | + except: pass |
| 84 | + return "N/A" |
| 85 | + |
| 86 | +def main(): |
| 87 | + parser = argparse.ArgumentParser(description="Aegis-AI Physical Model Profiler") |
| 88 | + parser.add_argument("--model", required=True, help="Model ID (e.g. gemma-4-26b-a4b-it-4bit)") |
| 89 | + parser.add_argument("--out", default="./profiling_results.md", help="Output markdown file path") |
| 90 | + args = parser.parse_args() |
| 91 | + |
| 92 | + results = [] |
| 93 | + subprocess.run(["killall", "SwiftLM"], stderr=subprocess.DEVNULL) |
| 94 | + |
| 95 | + for config in CONFIGS: |
| 96 | + print(f"\n--- Profiling {args.model} [{config['name']}] ---") |
| 97 | + model_path = f"/Users/simba/.aegis-ai/models/mlx_models/mlx-community/{args.model}" |
| 98 | + |
| 99 | + log_path = "./tmp/profile_server.log" |
| 100 | + cmd = [SWIFTLM_PATH, "--model", model_path] + config["flags"] |
| 101 | + |
| 102 | + with open(log_path, "w") as root_log: |
| 103 | + server_proc = subprocess.Popen(cmd, stdout=root_log, stderr=subprocess.STDOUT) |
| 104 | + |
| 105 | + if not poll_health(): |
| 106 | + print("Server failed to start.") |
| 107 | + server_proc.terminate() |
| 108 | + continue |
| 109 | + |
| 110 | + static_mem = extract_base_memory(log_path) |
| 111 | + |
| 112 | + print("Running 20-token test (prefill ~512, max ~20)...") |
| 113 | + ok, ttft, tps = make_request_stream(prompt_len=512, max_tokens=20) |
| 114 | + |
| 115 | + server_proc.send_signal(subprocess.signal.SIGTERM) |
| 116 | + server_proc.wait(timeout=10) |
| 117 | + |
| 118 | + real_mem = extract_real_memory(log_path) |
| 119 | + |
| 120 | + if ok: |
| 121 | + results.append({ |
| 122 | + "config": config["name"], |
| 123 | + "ttft_20": f"{ttft:.2f}", |
| 124 | + "tps_20": f"{tps:.2f}", |
| 125 | + "static_mem": static_mem, |
| 126 | + "real_mem": real_mem |
| 127 | + }) |
| 128 | + print(f"Result [{config['name']}]: TTFT={ttft:.2f}s TPS={tps:.2f} BaseRAM={static_mem} PhysRAM={real_mem}") |
| 129 | + |
| 130 | + with open(args.out, "w") as f: |
| 131 | + f.write(f"### `{args.model}` - Throughput & OS Memory Profile\n\n") |
| 132 | + f.write("| Configuration | Time To First Token | Generation Speed | Theoretical Reservation | Physical OS Footprint (RAM) |\n") |
| 133 | + f.write("|---|---|---|---|---|\n") |
| 134 | + for r in results: |
| 135 | + f.write(f"| {r['config']} | {r['ttft_20']}s | {r['tps_20']} tok/s | {r['static_mem']} | {r['real_mem']} |\n") |
| 136 | + |
| 137 | + print(f"\nDone. Results saved to {args.out}") |
| 138 | + |
| 139 | +if __name__ == "__main__": |
| 140 | + main() |
0 commit comments