#!/usr/bin/env python3 """VRAM/Context Detection Script for automaton. Detects GPU VRAM, system RAM, and model context window to recommend a safe context window for task decomposition. Usage: python vram_detect.py [model_name] [--model model_name] [--project project_dir] """ from __future__ import annotations import argparse import json import os import platform import re import shutil import subprocess import sys from pathlib import Path from typing import Optional # Known model context windows in tokens. MODEL_CONTEXT_WINDOWS: dict[str, int] = { "gpt-4o": 128_000, "gpt-4o-2024-05-13": 128_000, "gpt-4o-2024-08-06": 128_000, "gpt-4o-mini": 128_000, "gpt-4o-mini-2024-07-18": 128_000, "gpt-4-turbo": 128_000, "gpt-4-turbo-2024-04-09": 128_000, "gpt-4": 128_000, "gpt-4-0125-preview": 128_000, "gpt-4-1106-preview": 128_000, "claude-3-5-sonnet": 200_000, "claude-3-5-sonnet-20241022": 200_000, "claude-3-5-haiku": 200_000, "claude-3-5-haiku-20241022": 200_000, "claude-3-opus": 200_000, "claude-3-opus-20240229": 200_000, "claude-3-sonnet": 200_000, "claude-3-sonnet-20240229": 200_000, "claude-3-haiku": 200_000, "claude-3-haiku-20240307": 200_000, "claude-2": 200_000, "claude-2.1": 200_000, # Source: https://huggingface.co/meta-llama/Meta-Llama-3.1-8B (128k context, RoPE scaling) "llama-3.1-8b": 128_000, # Source: https://huggingface.co/meta-llama/Llama-3.3-70B-Instruct (128k context, RoPE scaling) "llama-3.3-70b": 128_000, # Source: https://huggingface.co/Qwen/Qwen2.5-7B-Instruct (128k context via YaRN) "qwen2.5-7b": 128_000, # Source: https://huggingface.co/Qwen/Qwen2.5-72B-Instruct (128k context via YaRN) "qwen2.5-72b": 128_000, # Source: https://huggingface.co/mistralai/Mistral-7B-Instruct-v0.3 (32k context) "mistral-7b": 32_000, # Source: https://huggingface.co/mistralai/Mistral-Large-Instruct-2411 (128k context) "mistral-large": 128_000, # Source: https://huggingface.co/deepseek-ai/DeepSeek-R1 (64k context, 128k claimed via YaRN; using conservative 64k) "deepseek-r1": 64_000, # Source: https://huggingface.co/deepseek-ai/DeepSeek-V3 (64k context, 128k claimed via YaRN; using conservative 64k) "deepseek-v3": 64_000, # Source: https://huggingface.co/THUDM/glm-4-9b-chat (128k context) "glm-4": 128_000, # Source: https://huggingface.co/THUDM/glm-4.5 (('128K-1M') context; using safe 128k) "glm-4.5": 128_000, # Source: https://huggingface.co/google/gemma-2-2b (8k context) "gemma-2": 8_000, # Source: https://huggingface.co/google/gemma-2-27b (8k context) "gemma-2-27b": 8_000, # Source: https://huggingface.co/microsoft/Phi-3-medium-128k-instruct (128k context) "phi-3": 128_000, # Source: https://huggingface.co/microsoft/phi-4 (16k context) "phi-4": 16_000, } DEFAULT_FALLBACK_CONTEXT_TOKENS = 128_000 DEFAULT_HEADROOM_PCT = 25 MAX_CONFIG_READ_BYTES = 10 * 1024 # 10KB limit per prompt requirement def run_command(cmd: list[str], timeout: float = 5.0) -> Optional[str]: """Run a command and return stdout, or None on failure.""" if platform.system() == 'Windows': first_arg = cmd[0] if cmd else '' if first_arg.startswith('Get-') or 'CimInstance' in first_arg: cmd = ['powershell', '-NoProfile', '-NoLogo', '-Command', ' '.join(cmd)] if not cmd or not shutil.which(cmd[0]): return None try: result = subprocess.run( cmd, capture_output=True, text=True, timeout=timeout, check=False, ) if result.returncode != 0: return None return result.stdout.strip() except (subprocess.TimeoutExpired, OSError): return None def detect_gpu_vram() -> tuple[int, int, int]: """Detect GPU VRAM and return (total_vram_kb, vram_per_gpu_kb, num_gpus).""" total_vram_kb = 0 num_gpus = 0 system = platform.system() if system == "Linux": nvidia_output = run_command( ["nvidia-smi", "--query-gpu=memory.total", "--format=csv,noheader,nounits"] ) if nvidia_output: lines = [line.strip() for line in nvidia_output.splitlines() if line.strip()] if lines: try: vram_mb = int(lines[0]) if vram_mb > 0: total_vram_kb = vram_mb * 1024 num_gpus = len(lines) print("GPU: NVIDIA (nvidia-smi available)") print(f"VRAM per GPU: {vram_mb // 1024}GB ({vram_mb} MB)") print(f"Num GPUs: {num_gpus}") except ValueError: print("GPU: NVIDIA (nvidia-smi available but driver not responding)") if total_vram_kb == 0: lspci_output = run_command(["lspci"]) if lspci_output and re.search(r"VGA|3D|Display", lspci_output, re.IGNORECASE): print("GPU detected via lspci") vram_mb = _detect_amd_vram_from_lspci() if vram_mb > 0: total_vram_kb = vram_mb * 1024 num_gpus = 1 print(f"VRAM: {vram_mb} MB ({vram_mb // 1024}GB)") elif system == "Darwin": sp_output = run_command(["system_profiler", "SPDisplaysDataType"]) if sp_output: if re.search(r"Chipset Model: Apple M\d", sp_output): mem_output = run_command(["sysctl", "-n", "hw.memsize"]) if mem_output: bytes_val = int(mem_output.strip()) total_vram_kb = (bytes_val // 1024) num_gpus = 1 print("GPU: Apple Silicon (unified memory)") print(f"Total System RAM (Shared VRAM): {total_vram_kb // 1024} MB") else: vram_match = re.search(r"VRAM \(Total\):\s*(\d+)\s*(GB|MB)", sp_output, re.IGNORECASE) if vram_match: val = int(vram_match.group(1)) unit = vram_match.group(2).upper() if unit == "GB": total_vram_kb = val * 1024 * 1024 else: total_vram_kb = val * 1024 num_gpus = 1 print("GPU: Intel Mac (dedicated VRAM)") print(f"VRAM: {val} {unit}") elif system == "Windows": wmic_output = run_command(["wmic", "path", "win32_VideoController", "get", "AdapterRAM,Name", "/format:list"]) if wmic_output: total_bytes = 0 count = 0 entries = re.split(r'(?=AdapterRAM=)', wmic_output) for entry in entries: if not entry.strip(): continue ram_match = re.search(r"AdapterRAM=(\d+)", entry) if ram_match: total_bytes += int(ram_match.group(1)) count += 1 if count > 0: total_vram_kb = total_bytes // 1024 num_gpus = count print(f"GPU: Windows (WMIC detected {count} GPUs)") print(f"Total VRAM: {total_vram_kb // 1024} MB") else: ps_output = run_command(["powershell", "-NoProfile", "-Command", "Get-CimInstance Win32_VideoController -Property AdapterRAM"]) if ps_output: ram_matches = re.findall(r"AdapterRAM=(\d+)", ps_output) if ram_matches: total_bytes = sum(int(x) for x in ram_matches) total_vram_kb = total_bytes // 1024 num_gpus = len(ram_matches) print(f"GPU: Windows (PowerShell fallback detected {num_gpus} GPUs)") print(f"Total VRAM: {total_vram_kb // 1024} MB") vram_per_gpu_kb = total_vram_kb // num_gpus if num_gpus > 0 else 0 return total_vram_kb, vram_per_gpu_kb, num_gpus def _detect_amd_vram_from_lspci() -> int: """Attempt to sum memory region sizes from lspci -vnn for VGA/3D/Display devices.""" output = run_command(["lspci", "-vnn"]) if not output: return 0 total_mb = 0 # Split into device blocks. Each block starts with a bus address like "71:00.0". blocks = re.split(r"\n\n", output) for block in blocks: if not re.search(r"VGA|3D|Display", block, re.IGNORECASE): continue for line in block.splitlines(): match = re.search(r"Memory at [0-9a-fA-Fx]+ \(.*\) \[size=(\d+)([MGK])\]", line) if match: value = int(match.group(1)) unit = match.group(2).upper() if unit == "G": total_mb += value * 1024 elif unit == "M": total_mb += value elif unit == "K": total_mb += value // 1024 return total_mb def _detect_ram_linux() -> tuple[int, int]: """Linux-only RAM detection via /proc/meminfo. Returns (total_kb, available_kb).""" meminfo = Path("/proc/meminfo") if meminfo.exists(): try: text = meminfo.read_text(encoding="utf-8") total_kb = _parse_meminfo_value(text, "MemTotal") available_kb = _parse_meminfo_value(text, "MemAvailable") or total_kb if total_kb > 0: print(f"RAM: {total_kb // 1024 // 1024}GB total, {available_kb // 1024 // 1024}GB available") return total_kb, available_kb except (OSError, ValueError): pass return 0, 0 def _detect_ram_darwin() -> tuple[int, int]: """macOS RAM detection via sysctl. Returns (total_kb, total_kb).""" try: sysctl_output = run_command(["sysctl", "-n", "hw.memsize"]) if sysctl_output: total_kb = int(sysctl_output) // 1024 if total_kb > 0: print("available RAM detection not supported on macOS, reporting total") print(f"RAM: {total_kb // 1024 // 1024}GB total") return total_kb, total_kb except (ValueError, OSError): pass return 0, 0 def _detect_ram_windows() -> tuple[int, int]: """Windows RAM detection via wmic / PowerShell. Returns (total_kb, total_kb).""" try: wmic_output = run_command(["wmic", "ComputerSystem", "get", "TotalPhysicalMemory", "/format:list"]) total_bytes = 0 if wmic_output: for line in wmic_output.splitlines(): if line.startswith("TotalPhysicalMemory="): total_bytes = int(line.split("=")[1]) break if total_bytes == 0: ps_output = run_command(["powershell", "-NoProfile", "-Command", "(Get-CimInstance Win32_ComputerSystem).TotalPhysicalMemory"]) if ps_output: total_bytes = int(ps_output.strip()) if total_bytes > 0: total_kb = total_bytes // 1024 print(f"RAM: {total_kb // 1024 // 1024}GB total, {total_kb // 1024 // 1024}GB available") return total_kb, total_kb except (ValueError, OSError): pass return 0, 0 def detect_ram() -> tuple[int, int]: """Detect total and available RAM in KB.""" system = platform.system() if system == "Linux": return _detect_ram_linux() elif system == "Darwin": return _detect_ram_darwin() elif system == "Windows": return _detect_ram_windows() print("RAM: Could not detect") return 0, 0 def _parse_meminfo_value(text: str, key: str) -> int: """Parse a value in KB from /proc/meminfo.""" for line in text.splitlines(): if line.startswith(key + ":"): parts = line.split() if len(parts) >= 2: return int(parts[1]) return 0 def _probe_ollama_model() -> Optional[str]: """Probe local ollama list for the first running model name. Returns None on any failure.""" try: if not shutil.which('ollama'): return None output = run_command(['ollama', 'list'], timeout=5.0) if not output: return None lines = output.strip().splitlines() data_rows = [line for line in lines if line.strip() and not line.startswith('NAME')] if not data_rows: return None first_row = data_rows[0] parts = first_row.split() if not parts: return None model_name = parts[0] if model_name.endswith(':latest'): model_name = model_name[:-len(':latest')] return model_name except Exception: return None def detect_model_context( model_name: Optional[str] = None, project_dir: Optional[Path] = None, ) -> int: """Detect model context window in tokens.""" if model_name: return _lookup_model_context(model_name) # Try config.md (global framework model settings). config_md = Path.home() / ".automaton" / "config.md" model_from_config, override_context = _parse_config_model(config_md) if model_from_config: print(f"Found model in config.md: {model_from_config}") if override_context and override_context != "auto": print(f"Using override context window from config.md: {override_context}") return override_context return _lookup_model_context(model_from_config) # Try .agent.md (project-level model override). if project_dir: agent_md = project_dir / ".automaton" / ".agent.md" model_from_agent, _ = _parse_config_model(agent_md) if model_from_agent: print(f"Found model in .agent.md: {model_from_agent}") if override_context and override_context != "auto": print(f"Using override context window from config.md: {override_context}") return override_context return _lookup_model_context(model_from_agent) # Try common API config files. config_files = [ ".env", ".env.local", "config.yaml", "config.yml", "config.json", "settings.yaml", ".automaton/config.yaml", ".automaton/config.json", ] if project_dir: for config_file in config_files: for candidate in [project_dir / config_file, project_dir / ".automaton" / config_file]: if candidate.exists(): model = _extract_model_from_file(candidate) if model: print(f"Found model in {candidate}: {model}") if override_context and override_context != "auto": print(f"Using override context window from config.md: {override_context}") return override_context return _lookup_model_context(model) # Try ollama probe (local LLMs via ollama). ollama_model = _probe_ollama_model() if ollama_model: print(f"Found model via ollama: {ollama_model}") return _lookup_model_context(ollama_model) print("Model: Unknown (could not detect from .agent.md or config files)") return 0 def _lookup_model_context(model_name: str) -> int: """Look up context window for a known model name.""" # Strip common version/date suffixes for lookup. for key in MODEL_CONTEXT_WINDOWS: if model_name.lower().startswith(key.lower()): print(f"Model: {model_name}") print(f"Context window: {MODEL_CONTEXT_WINDOWS[key] // 1000}k tokens") return MODEL_CONTEXT_WINDOWS[key] print(f"Model: {model_name} (unknown context window)") return 0 def _parse_config_model(config_path: Path) -> tuple[Optional[str], Optional[int]]: """Parse model name and override context window from a config file.""" if not config_path.exists(): return None, None try: text = config_path.read_text(encoding="utf-8", errors="replace") except OSError: return None, None model_name = None override_tokens: Optional[int] = None in_code_block = False for raw_line in text.splitlines(): line = raw_line.strip() if line.startswith("```"): in_code_block = not in_code_block continue if in_code_block: continue if line.startswith("#"): continue if re.search(r"^[-*]?\s*\*\*Model\*\*:|^\s*model\s*[:=]", line, re.IGNORECASE): if "override" in line.lower(): continue value = _extract_value(line) if value and value.lower() != "auto": model_name = value if re.search(r"^[-*]?\s*\*\*Override context( window)?\*\*:|^\s*override.*context\s*[:=]", line, re.IGNORECASE): value = _extract_value(line) if value and value.lower() != "auto": override_tokens = _parse_token_value(value) return model_name, override_tokens def _extract_model_from_file(path: Path) -> Optional[str]: """Extract model name from an API config file, reading at most 10KB.""" try: raw = path.read_bytes() text = raw[:MAX_CONFIG_READ_BYTES].decode("utf-8", errors="replace") except OSError: return None for line in text.splitlines(): line = line.strip() if line.startswith("#"): continue if re.search(r"model\s*[:=]", line, re.IGNORECASE): # Skip non-model keys like max_tokens, temperature, stream. key = line.split("=", 1)[0].split(":", 1)[0].strip().lower() if any(bad in key for bad in ["context", "max_tokens", "temperature", "stream"]): continue value = _extract_value(line) if value and value.lower() != "auto": return value return None def _extract_value(line: str) -> Optional[str]: """Extract the value after ':' or '=' from a key-value line.""" for sep in [":", "="]: if sep in line: value = line.split(sep, 1)[1].strip() value = value.split("#", 1)[0].strip() value = value.strip('"').strip("'") return value if value else None return None def _parse_token_value(value: str) -> Optional[int]: """Parse a token value like '128k', '5.6k', or '128000' into an integer.""" value = value.strip().lower() match = re.search(r"(\d+(?:\.\d+)?)\s*k?", value) if not match: return None number = float(match.group(1)) if "k" in value: return int(number * 1000) return int(number) def calculate_overhead(project_dir: Optional[Path] = None) -> int: """Estimate framework overhead in tokens.""" if project_dir: base_files = [ project_dir / ".automaton" / ".agent.md", project_dir / ".automaton" / ".rules.md", project_dir / ".automaton" / "prompts" / "workflow.md", project_dir / ".automaton" / "prompts" / "orchestrate.md", ] if not base_files[0].exists(): base_files = [ Path.home() / ".automaton" / ".agent.md", Path.home() / ".automaton" / ".rules.md", Path.home() / ".automaton" / "prompts" / "workflow.md", Path.home() / ".automaton" / "prompts" / "orchestrate.md", ] else: base_files = [ Path.home() / ".automaton" / ".agent.md", Path.home() / ".automaton" / ".rules.md", Path.home() / ".automaton" / "prompts" / "workflow.md", Path.home() / ".automaton" / "prompts" / "orchestrate.md", ] overhead_tokens = 0 for file_path in base_files: if file_path.exists(): try: chars = len(file_path.read_text(encoding="utf-8", errors="replace")) tokens = chars // 4 overhead_tokens += tokens print(f" {file_path.name}: ~{tokens} tokens") except OSError: pass print(f"Framework overhead: ~{overhead_tokens} tokens") return overhead_tokens def _iter_config_lines(text: str, section_header: str) -> list[str]: """Return non-code-block lines within a markdown section.""" lines: list[str] = [] in_section = False in_code_block = False for raw_line in text.splitlines(): line = raw_line.strip() if line.startswith("```"): in_code_block = not in_code_block continue if in_code_block: continue if line.startswith(section_header): in_section = True continue if in_section and line.startswith("##"): in_section = False if not in_section: continue lines.append(line) return lines def parse_vram_config(config_path: Path) -> dict[str, object]: """Parse VRAM configuration from config.md.""" defaults: dict[str, object] = { "auto_detect": True, "target_context_kb": 0, "headroom_pct": DEFAULT_HEADROOM_PCT, "max_peak_kb": 0, } if not config_path.exists(): return defaults try: text = config_path.read_text(encoding="utf-8", errors="replace") except OSError: return defaults for line in _iter_config_lines(text, "## VRAM Configuration"): if line.startswith("#"): continue lower = line.lower() if "auto-detect" in lower: value = _extract_value(line) if value: defaults["auto_detect"] = value.lower() in ("yes", "true", "1") elif "target" in lower and "context" in lower: value = _extract_value(line) if value: parsed = _parse_token_value(value) if parsed: defaults["target_context_kb"] = parsed elif "headroom" in lower: value = _extract_value(line) if value: match = re.search(r"(\d+)", value) if match: defaults["headroom_pct"] = int(match.group(1)) elif "max peak" in lower: value = _extract_value(line) if value: parsed = _parse_token_value(value) if parsed: defaults["max_peak_kb"] = parsed return defaults def recommend_context( gpu_vram_gb: int, ram_gb: int, model_context_kb: int, overhead_tokens: int, config: dict[str, object], ) -> tuple[int, int, int]: """Return (headroom_pct, recommended_kb, max_peak_kb).""" headroom_pct = int(config.get("headroom_pct", DEFAULT_HEADROOM_PCT)) # Manual override mode. if not config.get("auto_detect", True): target_kb = int(config.get("target_context_kb", 0)) max_peak_kb = int(config.get("max_peak_kb", 0)) if target_kb > 0: if max_peak_kb == 0 and headroom_pct > 0: max_peak_kb = target_kb * (100 - headroom_pct) // 100 return headroom_pct, target_kb, max_peak_kb recommended_kb = 0 if gpu_vram_gb >= 4: # Conservative: 1GB VRAM ≈ 2k context tokens. vram_context_kb = gpu_vram_gb * 2000 recommended_kb = vram_context_kb * (100 - headroom_pct) // 100 elif model_context_kb > 0: recommended_kb = model_context_kb * (100 - headroom_pct) // 100 else: # RAM fallback: 0.75k tokens per GB. ram_context_kb = ram_gb * 750 recommended_kb = ram_context_kb * (100 - headroom_pct) // 100 # Subtract framework overhead. net_kb = max(0, recommended_kb - overhead_tokens) # Calculate max peak context based on headroom. max_peak_kb = net_kb * (100 - headroom_pct) // 100 return headroom_pct, net_kb, max_peak_kb def main() -> int: parser = argparse.ArgumentParser(description="Detect VRAM/context for automaton") parser.add_argument("model", nargs="?", help="Model name") parser.add_argument("--model", "-m", dest="model_flag", help="Model name") parser.add_argument("--project", "-p", type=Path, help="Project directory") args = parser.parse_args() model_name = args.model_flag or args.model project_dir = args.project.resolve() if args.project else None print("=== VRAM / Context Detection ===\n") print("--- GPU VRAM ---") total_vram_kb, vram_per_gpu_kb, num_gpus = detect_gpu_vram() gpu_vram_gb = total_vram_kb // 1024 // 1024 print(f"Total VRAM: {gpu_vram_gb}GB") print(f"VRAM per GPU: {vram_per_gpu_kb // 1024 // 1024}GB") print(f"Num GPUs: {num_gpus}\n") print("--- RAM ---") ram_kb, _ = detect_ram() ram_gb = ram_kb // 1024 // 1024 print(f"RAM: {ram_gb}GB\n") print("--- Model Context Window ---") model_context_kb = detect_model_context(model_name, project_dir) print() print("--- Framework Overhead ---") overhead_tokens = calculate_overhead(project_dir) print() print("--- Recommendation ---") config_path = Path.home() / ".automaton" / "config.md" config = parse_vram_config(config_path) headroom_pct, recommended_kb, max_peak_kb = recommend_context( gpu_vram_gb, ram_gb, model_context_kb, overhead_tokens, config ) recommended_k = recommended_kb // 1000 if recommended_kb > 0 else 8 max_peak_k = max_peak_kb // 1000 if max_peak_kb > 0 else 6 print(f"Target context: {recommended_k}k tokens") print(f"Headroom: {headroom_pct}%") print(f"Max peak context per sub-task: {max_peak_k}k tokens") print("\n=== JSON Output ===") output = { "gpu_vram_gb": gpu_vram_gb, "ram_gb": ram_gb, "model_context_kb": model_context_kb, "framework_overhead_tokens": overhead_tokens, "recommended_kb": recommended_kb, "recommended_k": recommended_k, "headroom": headroom_pct / 100.0, "max_peak_context_kb": max_peak_kb, } print(json.dumps(output, indent=4)) return 0 if __name__ == "__main__": sys.exit(main())