#!/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 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, } 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 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 # Try nvidia-smi first. 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)") # Fallback: lspci for AMD/others. 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)") 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() -> tuple[int, int]: """Detect total and available RAM in 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 sysctl_output = run_command(["sysctl", "-n", "hw.memsize"]) if sysctl_output: try: total_kb = int(sysctl_output) // 1024 if total_kb > 0: print(f"RAM: {total_kb // 1024 // 1024}GB total") return total_kb, total_kb except ValueError: pass 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 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) 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())