Files
automaton/scripts/vram_detect.py
T

537 lines
18 KiB
Python
Raw Normal View History

#!/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())