CI / build (push) Has been cancelled
Batch 1 (High severity): - Bug 1: --audit cat3 now checks .automaton/tasks/ paths - Bug 4: Verdict PASS/FAIL uses structured ## Status: line parsing - Bug 5: register-guards.sh checks .json/.jsonc, writes plugin key, strips comments - Bug 7: --can-edit/--scope-check path prefix uses os.sep boundary Batch 2 (Medium/Low severity): - Bug 2: migrate-project.sh find command parentheses for -prune binding - Bug 3: vram_detect model prefix matching with known-suffix whitelist - Bug 6: dashboard reads .state file before artifact heuristic fallback - Bug 8: removed wildcard CORS, added security headers (nosniff, DENY) - Bug 9: stale-task detection uses .state.lastedit instead of .state mtime - Bug 10: TEST_PLAN.md maps to test_design (was implement) 249 tests pass (up from 235). All 10 tasks driven through full workflow to completion.
722 lines
27 KiB
Python
Executable File
722 lines
27 KiB
Python
Executable File
#!/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
|
|
|
|
|
|
_KNOWN_MODEL_SUFFIXES = {"instruct", "chat", "it", "fp16", "f16", "bf16"}
|
|
|
|
|
|
def _lookup_model_context(model_name: str) -> int:
|
|
"""Look up context window for a known model name.
|
|
|
|
Tries exact match first, then:
|
|
- ``key + ":"`` prefix (Ollama parameter tag, e.g. ``deepseek-r1:7b``)
|
|
- ``key + "-"`` prefix only if the next segment is a known instruction-tuning
|
|
suffix (e.g. ``llama-3.1-8b-instruct`` matches ``llama-3.1-8b``)
|
|
|
|
This prevents false matches like ``phi-4`` matching ``phi-4-mini-instruct``
|
|
or ``gpt-4o`` matching ``gpt-4o-foo-unknown``.
|
|
Longer keys are tried first so the most specific match wins.
|
|
"""
|
|
name_lower = model_name.lower()
|
|
for key in sorted(MODEL_CONTEXT_WINDOWS, key=len, reverse=True):
|
|
key_lower = key.lower()
|
|
if name_lower == key_lower:
|
|
print(f"Model: {model_name}")
|
|
print(f"Context window: {MODEL_CONTEXT_WINDOWS[key] // 1000}k tokens")
|
|
return MODEL_CONTEXT_WINDOWS[key]
|
|
if 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]
|
|
if name_lower.startswith(key_lower + "-"):
|
|
next_segment = name_lower[len(key_lower) + 1:].split("-")[0]
|
|
if next_segment in _KNOWN_MODEL_SUFFIXES:
|
|
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())
|