diff --git a/packages/ai/scripts/cursor-log.py b/packages/ai/scripts/cursor-log.py index ec1415815..ed92654f4 100755 --- a/packages/ai/scripts/cursor-log.py +++ b/packages/ai/scripts/cursor-log.py @@ -18,232 +18,260 @@ from pathlib import Path SKIP_DELTAS = {"tokenDelta", "partialToolCall", "heartbeat", "thinkingDelta"} COALESCE_DELTAS = {"textDelta"} + def get_delta_type(entry: dict) -> str | None: - """Return delta type if this is a delta entry, else None.""" - typ = entry.get("type", "") - subtype = entry.get("subtype", "") + """Return delta type if this is a delta entry, else None.""" + typ = entry.get("type", "") + subtype = entry.get("subtype", "") - if typ == "interactionUpdate" and subtype in SKIP_DELTAS | COALESCE_DELTAS: - return subtype + if typ == "interactionUpdate" and subtype in SKIP_DELTAS | COALESCE_DELTAS: + return subtype - if typ == "info" and subtype == "interactionUpdate": - data = entry.get("data", {}) - if isinstance(data, dict): - return data.get("updateCase") + if typ == "info" and subtype == "interactionUpdate": + data = entry.get("data", {}) + if isinstance(data, dict): + return data.get("updateCase") + + return None - return None def is_noise(entry: dict, verbose: bool = False) -> bool: - if verbose: - return False + if verbose: + return False - typ = entry.get("type", "") - subtype = entry.get("subtype", "") + typ = entry.get("type", "") + subtype = entry.get("subtype", "") - # serverMessage:interactionUpdate is redundant (we log the inner update) - if typ == "serverMessage" and subtype == "interactionUpdate": - return True + # serverMessage:interactionUpdate is redundant (we log the inner update) + if typ == "serverMessage" and subtype == "interactionUpdate": + return True - # KV blob ops are noisy - if typ == "kvClient" or (typ == "serverMessage" and subtype == "kvServerMessage"): - return True + # KV blob ops are noisy + if typ == "kvClient" or (typ == "serverMessage" and subtype == "kvServerMessage"): + return True - # conversationCheckpointUpdate is noisy - if typ == "serverMessage" and subtype == "conversationCheckpointUpdate": - return True + # conversationCheckpointUpdate is noisy + if typ == "serverMessage" and subtype == "conversationCheckpointUpdate": + return True - # execClient:requestContextResult is redundant with info:execClientMessage - if typ == "execClient" and subtype == "requestContextResult": - return True + # execClient:requestContextResult is redundant with info:execClientMessage + if typ == "execClient" and subtype == "requestContextResult": + return True - # Filter streaming deltas that we skip entirely - delta_type = get_delta_type(entry) - if delta_type in SKIP_DELTAS: - return True + # Filter streaming deltas that we skip entirely + delta_type = get_delta_type(entry) + if delta_type in SKIP_DELTAS: + return True + + return False - return False def format_data(typ: str, subtype: str, data: dict | None) -> str: - if not data or not isinstance(data, dict): - return "" + if not data or not isinstance(data, dict): + return "" - # Extract useful fields based on message type - if typ == "serverMessage" and subtype == "execServerMessage": - msg = data.get("message", {}) - case = msg.get("case", "") - value = msg.get("value", {}) - if case == "mcpArgs": - name = value.get("name") or value.get("toolName") or "?" - args = value.get("args", {}) - args_str = json.dumps(args, default=str) if args else "" - if len(args_str) > 200: - args_str = args_str[:200] + "..." - return f" mcp:{name} {args_str}" - elif case == "grepArgs": - pattern = value.get("pattern", "") - path = value.get("path", ".") - return f" grep:{pattern[:30]}@{path[:30]}" if pattern else f" grep@{path[:30]}" - elif case == "shellArgs": - cmd = value.get("command", "")[:50] - return f" shell:{cmd}" - elif case in ("readArgs", "writeArgs", "lsArgs", "deleteArgs"): - path = value.get("path", "") - return f" {case.replace('Args', '')}:{path[:60]}" if path else f" {case}" - return f" {case}" + # Extract useful fields based on message type + if typ == "serverMessage" and subtype == "execServerMessage": + msg = data.get("message", {}) + case = msg.get("case", "") + value = msg.get("value", {}) + if case == "mcpArgs": + name = value.get("name") or value.get("toolName") or "?" + args = value.get("args", {}) + args_str = json.dumps(args, default=str) if args else "" + if len(args_str) > 200: + args_str = args_str[:200] + "..." + return f" mcp:{name} {args_str}" + elif case == "grepArgs": + pattern = value.get("pattern", "") + path = value.get("path", ".") + return ( + f" grep:{pattern[:30]}@{path[:30]}" if pattern else f" grep@{path[:30]}" + ) + elif case == "shellArgs": + cmd = value.get("command", "")[:50] + return f" shell:{cmd}" + elif case in ("readArgs", "writeArgs", "lsArgs", "deleteArgs"): + path = value.get("path", "") + return f" {case.replace('Args', '')}:{path[:60]}" if path else f" {case}" + return f" {case}" - if typ == "info" and subtype == "interactionUpdate": - update_case = data.get("updateCase", "") - return f" {update_case}" + if typ == "info" and subtype == "interactionUpdate": + update_case = data.get("updateCase", "") + return f" {update_case}" - if typ == "info" and subtype == "builtRunRequest": - return f" tools={data.get('tools', 0)}" + if typ == "info" and subtype == "builtRunRequest": + return f" tools={data.get('tools', 0)}" - if typ == "info" and subtype == "execClientMessage": - return f" {data.get('messageCase', '')}" + if typ == "info" and subtype == "execClientMessage": + return f" {data.get('messageCase', '')}" + + # Default: show compact JSON + filtered = { + k: v + for k, v in data.items() + if v is not None and k not in ("detail", "$typeName") + } + if not filtered: + return "" + s = json.dumps(filtered, default=str) + return f" {s[:120]}..." if len(s) > 120 else f" {s}" - # Default: show compact JSON - filtered = {k: v for k, v in data.items() if v is not None and k not in ("detail", "$typeName")} - if not filtered: - return "" - s = json.dumps(filtered, default=str) - return f" {s[:120]}..." if len(s) > 120 else f" {s}" def format_entry(entry: dict, verbose: bool = False) -> str | None: - if is_noise(entry, verbose): - return None + if is_noise(entry, verbose): + return None - ts = entry.get("ts", 0) - typ = entry.get("type", "?") - subtype = entry.get("subtype") - data = entry.get("data") + ts = entry.get("ts", 0) + typ = entry.get("type", "?") + subtype = entry.get("subtype") + data = entry.get("data") - time_str = datetime.fromtimestamp(ts / 1000).strftime("%H:%M:%S.%f")[:-3] if ts else "??:??:??" + time_str = ( + datetime.fromtimestamp(ts / 1000).strftime("%H:%M:%S.%f")[:-3] + if ts + else "??:??:??" + ) - type_str = f"{typ}:{subtype}" if subtype else typ - data_str = format_data(typ, subtype or "", data) if not verbose else "" + type_str = f"{typ}:{subtype}" if subtype else typ + data_str = format_data(typ, subtype or "", data) if not verbose else "" - if verbose and data: - data_str = " " + json.dumps(data, default=str)[:300] + if verbose and data: + data_str = " " + json.dumps(data, default=str)[:300] + + return f"[{time_str}] {type_str}{data_str}" - return f"[{time_str}] {type_str}{data_str}" def extract_text_delta(entry: dict) -> str | None: - """Extract text from a textDelta entry.""" - data = entry.get("data", {}) - if isinstance(data, dict): - # Direct textDelta - if "text" in data: - return data["text"] - # Nested in message.value - msg = data.get("message", {}) - if isinstance(msg, dict): - value = msg.get("value", {}) - if isinstance(value, dict) and "text" in value: - return value["text"] - return None + """Extract text from a textDelta entry.""" + data = entry.get("data", {}) + if isinstance(data, dict): + # Direct textDelta + if "text" in data: + return data["text"] + # Nested in message.value + msg = data.get("message", {}) + if isinstance(msg, dict): + value = msg.get("value", {}) + if isinstance(value, dict) and "text" in value: + return value["text"] + return None + def coalesce_entries(entries: list[dict], verbose: bool = False) -> list[str]: - """Process entries, coalescing consecutive textDeltas.""" - output = [] - text_buffer = "" - text_ts = 0 + """Process entries, coalescing consecutive textDeltas.""" + output = [] + text_buffer = "" + text_ts = 0 - def flush_text(): - nonlocal text_buffer, text_ts - if text_buffer: - time_str = datetime.fromtimestamp(text_ts / 1000).strftime("%H:%M:%S.%f")[:-3] - text = text_buffer.replace("\n", "\\n") - if len(text) > 300: - text = text[:300] + "..." - output.append(f"[{time_str}] text: {text}") - text_buffer = "" - text_ts = 0 + def flush_text(): + nonlocal text_buffer, text_ts + if text_buffer: + time_str = datetime.fromtimestamp(text_ts / 1000).strftime("%H:%M:%S.%f")[ + :-3 + ] + text = text_buffer.replace("\n", "\\n") + if len(text) > 300: + text = text[:300] + "..." + output.append(f"[{time_str}] text: {text}") + text_buffer = "" + text_ts = 0 - for entry in entries: - typ = entry.get("type", "") - subtype = entry.get("subtype", "") + for entry in entries: + typ = entry.get("type", "") + subtype = entry.get("subtype", "") - # Accumulate textDelta - if typ == "interactionUpdate" and subtype == "textDelta": - text = extract_text_delta(entry) - if text: - if not text_buffer: - text_ts = entry.get("ts", 0) - text_buffer += text - continue + # Accumulate textDelta + if typ == "interactionUpdate" and subtype == "textDelta": + text = extract_text_delta(entry) + if text: + if not text_buffer: + text_ts = entry.get("ts", 0) + text_buffer += text + continue - # Skip noise entirely (don't flush for these) - if is_noise(entry, verbose): - continue + # Skip noise entirely (don't flush for these) + if is_noise(entry, verbose): + continue - # Real entry - flush text buffer first - flush_text() + # Real entry - flush text buffer first + flush_text() - formatted = format_entry(entry, verbose) - if formatted: - output.append(formatted) + formatted = format_entry(entry, verbose) + if formatted: + output.append(formatted) + + flush_text() + return output - flush_text() - return output def parse_entries(path: Path, last: int = 0) -> list[dict]: - """Parse JSONL file into entries.""" - entries = [] - lines = path.read_text().strip().split("\n") - if last > 0: - lines = lines[-last:] if len(lines) > last else lines - for line in lines: - if line.strip(): - try: - entries.append(json.loads(line)) - except json.JSONDecodeError: - print(f"[PARSE ERROR] {line[:100]}", file=sys.stderr) - return entries + """Parse JSONL file into entries.""" + entries = [] + lines = path.read_text().strip().split("\n") + if last > 0: + lines = lines[-last:] if len(lines) > last else lines + for line in lines: + if line.strip(): + try: + entries.append(json.loads(line)) + except json.JSONDecodeError: + print(f"[PARSE ERROR] {line[:100]}", file=sys.stderr) + return entries -def process_file(path: Path, verbose: bool = False, follow: bool = False, last: int = 0): - if not path.exists(): - print(f"File not found: {path}", file=sys.stderr) - sys.exit(1) - if follow: - # Follow mode: buffer briefly then emit - with open(path) as f: - f.seek(0, 2) - buffer = [] - last_emit = time.time() - while True: - line = f.readline() - if line and line.strip(): - try: - buffer.append(json.loads(line)) - except json.JSONDecodeError: - pass - # Emit buffered entries every 0.5s or when buffer is large - if buffer and (time.time() - last_emit > 0.5 or len(buffer) > 50): - for out in coalesce_entries(buffer, verbose): - print(out, flush=True) - buffer = [] - last_emit = time.time() - elif not line: - time.sleep(0.05) - else: - entries = parse_entries(path, last) - for out in coalesce_entries(entries, verbose): - print(out) +def process_file( + path: Path, verbose: bool = False, follow: bool = False, last: int = 0 +): + if not path.exists(): + print(f"File not found: {path}", file=sys.stderr) + sys.exit(1) + + if follow: + # Follow mode: buffer briefly then emit + with open(path) as f: + f.seek(0, 2) + buffer = [] + last_emit = time.time() + while True: + line = f.readline() + if line and line.strip(): + try: + buffer.append(json.loads(line)) + except json.JSONDecodeError: + pass + # Emit buffered entries every 0.5s or when buffer is large + if buffer and (time.time() - last_emit > 0.5 or len(buffer) > 50): + for out in coalesce_entries(buffer, verbose): + print(out, flush=True) + buffer = [] + last_emit = time.time() + elif not line: + time.sleep(0.05) + else: + entries = parse_entries(path, last) + for out in coalesce_entries(entries, verbose): + print(out) + def main(): - parser = argparse.ArgumentParser(description="Filter Cursor debug logs") - parser.add_argument("file", type=Path, help="JSONL log file") - parser.add_argument("-v", "--verbose", action="store_true", help="Show all entries") - parser.add_argument("-f", "--follow", action="store_true", help="Follow mode (tail -f)") - parser.add_argument("--last", type=int, default=0, help="Show last N entries") + parser = argparse.ArgumentParser(description="Filter Cursor debug logs") + parser.add_argument("file", type=Path, help="JSONL log file") + parser.add_argument("-v", "--verbose", action="store_true", help="Show all entries") + parser.add_argument( + "-f", "--follow", action="store_true", help="Follow mode (tail -f)" + ) + parser.add_argument("--last", type=int, default=0, help="Show last N entries") - args = parser.parse_args() + args = parser.parse_args() + + try: + process_file( + args.file, verbose=args.verbose, follow=args.follow, last=args.last + ) + except KeyboardInterrupt: + pass - try: - process_file(args.file, verbose=args.verbose, follow=args.follow, last=args.last) - except KeyboardInterrupt: - pass if __name__ == "__main__": - main() + main() diff --git a/packages/ai/scripts/proto-extractor.py b/packages/ai/scripts/proto-extractor.py index 39199f8a7..9a59bfbf3 100644 --- a/packages/ai/scripts/proto-extractor.py +++ b/packages/ai/scripts/proto-extractor.py @@ -13,17 +13,36 @@ from dataclasses import dataclass, field from typing import Optional SCALAR_TYPES = { - 1: "double", 2: "float", 3: "int64", 4: "uint64", 5: "int32", - 6: "fixed64", 7: "fixed32", 8: "bool", 9: "string", 10: "group", - 11: "message", 12: "bytes", 13: "uint32", 14: "enum", - 15: "sfixed32", 16: "sfixed64", 17: "sint32", 18: "sint64", + 1: "double", + 2: "float", + 3: "int64", + 4: "uint64", + 5: "int32", + 6: "fixed64", + 7: "fixed32", + 8: "bool", + 9: "string", + 10: "group", + 11: "message", + 12: "bytes", + 13: "uint32", + 14: "enum", + 15: "sfixed32", + 16: "sfixed64", + 17: "sint32", + 18: "sint64", } WEBPACK_NOISE = [ - '__webpack_require__', 'harmony export', 'harmony import', - 'use strict', 'WEBPACK_IMPORTED_MODULE', 'binding', + "__webpack_require__", + "harmony export", + "harmony import", + "use strict", + "WEBPACK_IMPORTED_MODULE", + "binding", ] + @dataclass class FieldDef: no: int @@ -36,24 +55,28 @@ class FieldDef: oneof: Optional[str] = None map_key: Optional[int] = None + @dataclass class MessageDef: type_name: str fields: list[FieldDef] = field(default_factory=list) comment: str = "" + @dataclass class EnumValueDef: name: str no: int comment: str = "" + @dataclass class EnumDef: type_name: str values: list[EnumValueDef] = field(default_factory=list) comment: str = "" + @dataclass class MethodDef: name: str @@ -62,12 +85,14 @@ class MethodDef: kind: str comment: str = "" + @dataclass class ServiceDef: type_name: str methods: list[MethodDef] = field(default_factory=list) comment: str = "" + @dataclass class ProtoFile: path: str @@ -90,24 +115,24 @@ def extract_jsdoc_comment(text: str) -> str: return "" lines = [] - for line in text.split('\n'): + for line in text.split("\n"): line = line.strip() - if line.startswith('*'): + if line.startswith("*"): line = line[1:].strip() - if line.startswith('/**') or line.endswith('*/'): + if line.startswith("/**") or line.endswith("*/"): continue - if '@generated' in line: + if "@generated" in line: continue if is_webpack_noise(line): continue - if line.startswith('case:') or '= { case:' in line: + if line.startswith("case:") or "= { case:" in line: continue - if line == '/' or line == '//' or len(line) <= 2: + if line == "/" or line == "//" or len(line) <= 2: continue if line: lines.append(line) - result = ' '.join(lines) + result = " ".join(lines) if is_webpack_noise(result): return "" if len(result) <= 2: @@ -117,16 +142,16 @@ def extract_jsdoc_comment(text: str) -> str: def parse_type_reference(rest: str) -> str: """Parse the T: field to extract the type reference.""" - webpack_comment = re.search(r'T:\s*[^,]+/\*\s*\.?(\w+)\s*\*/', rest) + webpack_comment = re.search(r"T:\s*[^,]+/\*\s*\.?(\w+)\s*\*/", rest) if webpack_comment: return webpack_comment.group(1) - type_match = re.search(r'T:\s*(\d+|[A-Za-z_][A-Za-z0-9_]*)', rest) + type_match = re.search(r"T:\s*(\d+|[A-Za-z_][A-Za-z0-9_]*)", rest) if type_match: t = type_match.group(1) if t.isdigit(): return SCALAR_TYPES.get(int(t), f"scalar_{t}") - if is_webpack_noise(t) or t.startswith('_'): + if is_webpack_noise(t) or t.startswith("_"): return "unknown" return t @@ -136,23 +161,31 @@ def parse_type_reference(rest: str) -> str: def parse_field_list(fields_str: str, field_comments: dict[str, str]) -> list[FieldDef]: """Parse the fields array from newFieldList.""" fields = [] - pattern = r'\{\s*no:\s*(\d+)\s*,\s*name:\s*"([^"]+)"\s*,\s*kind:\s*"([^"]+)"([^}]*)\}' + pattern = ( + r'\{\s*no:\s*(\d+)\s*,\s*name:\s*"([^"]+)"\s*,\s*kind:\s*"([^"]+)"([^}]*)\}' + ) for m in re.finditer(pattern, fields_str): no, name, kind, rest = int(m.group(1)), m.group(2), m.group(3), m.group(4) type_ref = parse_type_reference(rest) - fields.append(FieldDef( - no=no, - name=name, - kind=kind, - type_ref=type_ref, - comment=field_comments.get(name, ""), - opt='opt: true' in rest, - repeated='repeated: true' in rest, - oneof=m2.group(1) if (m2 := re.search(r'oneof:\s*"([^"]+)"', rest)) else None, - map_key=int(m2.group(1)) if (m2 := re.search(r'mapKey:\s*(\d+)', rest)) else None, - )) + fields.append( + FieldDef( + no=no, + name=name, + kind=kind, + type_ref=type_ref, + comment=field_comments.get(name, ""), + opt="opt: true" in rest, + repeated="repeated: true" in rest, + oneof=m2.group(1) + if (m2 := re.search(r'oneof:\s*"([^"]+)"', rest)) + else None, + map_key=int(m2.group(1)) + if (m2 := re.search(r"mapKey:\s*(\d+)", rest)) + else None, + ) + ) return fields @@ -160,19 +193,23 @@ def parse_field_list(fields_str: str, field_comments: dict[str, str]) -> list[Fi def extract_field_comments(class_body: str) -> dict[str, str]: """Extract field comments from class property declarations.""" comments = {} - pattern = r'/\*\*([\s\S]*?)@generated from field:[^*]*\*/\s*\n?\s*(\w+)' + pattern = r"/\*\*([\s\S]*?)@generated from field:[^*]*\*/\s*\n?\s*(\w+)" for m in re.finditer(pattern, class_body): comment_text = extract_jsdoc_comment(m.group(1)) field_name_camel = m.group(2) - field_name_snake = re.sub(r'([A-Z])', r'_\1', field_name_camel).lower().lstrip('_') + field_name_snake = ( + re.sub(r"([A-Z])", r"_\1", field_name_camel).lower().lstrip("_") + ) if comment_text and not is_webpack_noise(comment_text): comments[field_name_snake] = comment_text return comments -def find_file_for_pos(file_ranges: dict[str, list[tuple[int, int]]], pos: int) -> str | None: +def find_file_for_pos( + file_ranges: dict[str, list[tuple[int, int]]], pos: int +) -> str | None: """Find which file a position belongs to.""" for fp, segments in file_ranges.items(): for start, end in segments: @@ -181,19 +218,21 @@ def find_file_for_pos(file_ranges: dict[str, list[tuple[int, int]]], pos: int) - return None -def extract_messages(content: str, file_ranges: dict[str, list[tuple[int, int]]]) -> dict[str, list[MessageDef]]: +def extract_messages( + content: str, file_ranges: dict[str, list[tuple[int, int]]] +) -> dict[str, list[MessageDef]]: """Extract all message definitions grouped by file.""" messages_by_file: dict[str, list[MessageDef]] = {} seen_types: set[str] = set() pattern = re.compile( - r'/\*\*([\s\S]*?)@generated from message ([^\s*]+)[\s\S]*?\*/' - r'[\s\S]*?class (\w+) extends [^{]+\{' - r'([\s\S]*?)' + r"/\*\*([\s\S]*?)@generated from message ([^\s*]+)[\s\S]*?\*/" + r"[\s\S]*?class (\w+) extends [^{]+\{" + r"([\s\S]*?)" r'static typeName\s*=\s*"([^"]+)"' - r'[\s\S]*?' - r'static fields\s*=\s*[^(]+\(\(\)\s*=>\s*\[([\s\S]*?)\]\)', - re.MULTILINE + r"[\s\S]*?" + r"static fields\s*=\s*[^(]+\(\(\)\s*=>\s*\[([\s\S]*?)\]\)", + re.MULTILINE, ) for m in pattern.finditer(content): @@ -218,16 +257,18 @@ def extract_messages(content: str, file_ranges: dict[str, list[tuple[int, int]]] return messages_by_file -def extract_enums(content: str, file_ranges: dict[str, list[tuple[int, int]]]) -> dict[str, list[EnumDef]]: +def extract_enums( + content: str, file_ranges: dict[str, list[tuple[int, int]]] +) -> dict[str, list[EnumDef]]: """Extract all enum definitions grouped by file.""" enums_by_file: dict[str, list[EnumDef]] = {} seen_types: set[str] = set() enum_block_pattern = re.compile( - r'/\*\*([\s\S]*?)@generated from enum ([^\s*]+)[\s\S]*?\*/' - r'[\s\S]*?' + r"/\*\*([\s\S]*?)@generated from enum ([^\s*]+)[\s\S]*?\*/" + r"[\s\S]*?" r'setEnumType\([^,]+,\s*"([^"]+)",\s*\[([\s\S]*?)\]\)', - re.MULTILINE + re.MULTILINE, ) for m in enum_block_pattern.finditer(content): @@ -257,24 +298,26 @@ def extract_enums(content: str, file_ranges: dict[str, list[tuple[int, int]]]) - def resolve_webpack_type(ref: str) -> str: """Resolve a webpack type reference like 'agent_service_pb/* AgentClientMessage */.KS'.""" - webpack_comment = re.search(r'/\*\s*(\w+)\s*\*/', ref) + webpack_comment = re.search(r"/\*\s*(\w+)\s*\*/", ref) if webpack_comment: return webpack_comment.group(1) - parts = ref.replace(',', '').strip().split('.') + parts = ref.replace(",", "").strip().split(".") return parts[-1] if parts else ref -def extract_services(content: str, file_ranges: dict[str, list[tuple[int, int]]]) -> dict[str, list[ServiceDef]]: +def extract_services( + content: str, file_ranges: dict[str, list[tuple[int, int]]] +) -> dict[str, list[ServiceDef]]: """Extract all service definitions grouped by file.""" services_by_file: dict[str, list[ServiceDef]] = {} seen_types: set[str] = set() service_pattern = re.compile( - r'/\*\*([^*]|\*[^/])*@generated from service ([^\s*]+)([^*]|\*[^/])*\*/' - r'\s*(?:const|var)\s+\w+\s*=\s*\{' + r"/\*\*([^*]|\*[^/])*@generated from service ([^\s*]+)([^*]|\*[^/])*\*/" + r"\s*(?:const|var)\s+\w+\s*=\s*\{" r'[^}]*typeName:\s*"([^"]+)"' - r'[^}]*methods:\s*\{([\s\S]*?)\}\s*\}', - re.MULTILINE + r"[^}]*methods:\s*\{([\s\S]*?)\}\s*\}", + re.MULTILINE, ) for m in service_pattern.finditer(content): @@ -293,31 +336,33 @@ def extract_services(content: str, file_ranges: dict[str, list[tuple[int, int]]] # Extract comment from the match text before @generated full_match = m.group(0) - comment_end = full_match.find('@generated') + comment_end = full_match.find("@generated") comment_text = full_match[3:comment_end] if comment_end > 0 else "" comment = extract_jsdoc_comment(comment_text) methods = [] method_pattern = re.compile( - r'/\*\*([\s\S]*?)@generated from rpc [^\s*]+\.(\w+)[\s\S]*?\*/' - r'\s*\w+:\s*\{' + r"/\*\*([\s\S]*?)@generated from rpc [^\s*]+\.(\w+)[\s\S]*?\*/" + r"\s*\w+:\s*\{" r'[^}]*name:\s*"([^"]+)"' - r'[^}]*I:\s*([^,]+),' - r'[^}]*O:\s*([^,]+),' - r'[^}]*kind:\s*[^.]+\.(\w+)', - re.MULTILINE + r"[^}]*I:\s*([^,]+)," + r"[^}]*O:\s*([^,]+)," + r"[^}]*kind:\s*[^.]+\.(\w+)", + re.MULTILINE, ) for mm in method_pattern.finditer(methods_str): method_comment, _, name, input_ref, output_ref, kind = mm.groups() - methods.append(MethodDef( - name=name, - input_type=resolve_webpack_type(input_ref), - output_type=resolve_webpack_type(output_ref), - kind=kind, - comment=extract_jsdoc_comment(method_comment), - )) + methods.append( + MethodDef( + name=name, + input_type=resolve_webpack_type(input_ref), + output_type=resolve_webpack_type(output_ref), + kind=kind, + comment=extract_jsdoc_comment(method_comment), + ) + ) svc = ServiceDef(type_name=type_name, methods=methods, comment=comment) services_by_file.setdefault(file_path, []).append(svc) @@ -332,7 +377,7 @@ def find_file_ranges(content: str) -> dict[str, tuple[int, int]]: We collect all ranges and merge them so all occurrences are captured. """ file_pattern = re.compile( - r'// @generated from file ([^\s]+) \(package ([^,]+), syntax (\w+)\)' + r"// @generated from file ([^\s]+) \(package ([^,]+), syntax (\w+)\)" ) matches = list(file_pattern.finditer(content)) @@ -373,7 +418,7 @@ def field_to_proto(f: FieldDef, indent: str = " ") -> str: prefix = "repeated " lines.append(f"{indent}{prefix}{f.type_ref} {f.name} = {f.no};") - return '\n'.join(lines) + return "\n".join(lines) def get_simple_name(type_name: str) -> str: @@ -381,10 +426,10 @@ def get_simple_name(type_name: str) -> str: Handles nested types like 'agent.v1.Outer.Inner' by converting to 'Outer_Inner'. """ - parts = type_name.split('.') + parts = type_name.split(".") # Skip the package prefix (e.g., 'agent.v1') if len(parts) > 2: - return '_'.join(parts[2:]) + return "_".join(parts[2:]) return parts[-1] @@ -419,7 +464,7 @@ def message_to_proto(msg: MessageDef, indent: str = "") -> str: lines.append(f"{indent} }}") lines.append(f"{indent}}}") - return '\n'.join(lines) + return "\n".join(lines) def enum_to_proto(enum: EnumDef, indent: str = "") -> str: @@ -436,7 +481,7 @@ def enum_to_proto(enum: EnumDef, indent: str = "") -> str: lines.append(f"{indent} // {v.comment}") lines.append(f"{indent} {v.name} = {v.no};") lines.append(f"{indent}}}") - return '\n'.join(lines) + return "\n".join(lines) def service_to_proto(svc: ServiceDef, indent: str = "") -> str: @@ -456,34 +501,36 @@ def service_to_proto(svc: ServiceDef, indent: str = "") -> str: stream_in = "stream " if m.kind in ("ClientStreaming", "BiDiStreaming") else "" stream_out = "stream " if m.kind in ("ServerStreaming", "BiDiStreaming") else "" - lines.append(f"{indent} rpc {m.name}({stream_in}{m.input_type}) returns ({stream_out}{m.output_type});") + lines.append( + f"{indent} rpc {m.name}({stream_in}{m.input_type}) returns ({stream_out}{m.output_type});" + ) lines.append(f"{indent}}}") - return '\n'.join(lines) + return "\n".join(lines) def generate_proto_file(proto: ProtoFile) -> str: """Generate complete proto file content.""" lines = [ f'syntax = "{proto.syntax}";', - '', - f'package {proto.package};', - '', + "", + f"package {proto.package};", + "", ] for enum in proto.enums: lines.append(enum_to_proto(enum)) - lines.append('') + lines.append("") for msg in proto.messages: lines.append(message_to_proto(msg)) - lines.append('') + lines.append("") for svc in proto.services: lines.append(service_to_proto(svc)) - lines.append('') + lines.append("") - return '\n'.join(lines) + return "\n".join(lines) def main(): @@ -499,11 +546,11 @@ def main(): filter_pkg = None for arg in sys.argv[3:]: - if arg.startswith('--filter='): - filter_pkg = arg.split('=')[1] + if arg.startswith("--filter="): + filter_pkg = arg.split("=")[1] print(f"Reading {input_file}...", file=sys.stderr) - with open(input_file, 'r', encoding='utf-8', errors='replace') as f: + with open(input_file, "r", encoding="utf-8", errors="replace") as f: content = f.read() print(f"File size: {len(content) / 1024 / 1024:.2f} MB", file=sys.stderr) @@ -522,7 +569,7 @@ def main(): services_by_file = extract_services(content, file_ranges) file_pattern = re.compile( - r'// @generated from file ([^\s]+) \(package ([^,]+), syntax (\w+)\)' + r"// @generated from file ([^\s]+) \(package ([^,]+), syntax (\w+)\)" ) # Collect all messages, enums, services into one consolidated proto @@ -569,7 +616,10 @@ def main(): proto_content = generate_proto_file(consolidated) output_file.write_text(proto_content) - print(f"\nTotal: {len(all_messages)} messages, {len(all_enums)} enums, {len(all_services)} services", file=sys.stderr) + print( + f"\nTotal: {len(all_messages)} messages, {len(all_enums)} enums, {len(all_services)} services", + file=sys.stderr, + ) print(f"Output written to: {output_file}", file=sys.stderr) diff --git a/packages/coding-agent/src/eval/py/prelude.py b/packages/coding-agent/src/eval/py/prelude.py index bb84253fb..4feab9751 100644 --- a/packages/coding-agent/src/eval/py/prelude.py +++ b/packages/coding-agent/src/eval/py/prelude.py @@ -1,10 +1,12 @@ from __future__ import annotations + # OMP prelude helpers (loaded once into the runner namespace) if "__omp_prelude_loaded__" not in globals(): __omp_prelude_loaded__ = True from pathlib import Path import os, json, math, re from urllib.parse import unquote + INTENT_FIELD = "i" # __omp_display is injected by runner.py before the prelude executes; it @@ -40,7 +42,6 @@ if "__omp_prelude_loaded__" not in globals(): """Emit structured status event for TUI rendering.""" _omp_display({"application/x-omp-status": {"op": op, **data}}, raw=True) - def env(key: str | None = None, value: str | None = None): """Get/set environment variables.""" if key is None: @@ -150,18 +151,18 @@ if "__omp_prelude_loaded__" not in globals(): limit: int | None = None, ) -> str | dict | list[dict]: """Read task/agent output by ID. Returns text or JSON depending on format. - + Args: *ids: Output IDs to read (e.g., 'explore_0', 'reviewer_1') format: 'raw' (default), 'json' (dict with metadata), 'stripped' (no ANSI) query: jq-like query for JSON outputs (e.g., '.endpoints[0].file') offset: Line number to start reading from (1-indexed) limit: Maximum number of lines to read - + Returns: Single ID: str (format='raw'/'stripped') or dict (format='json') Multiple IDs: list of dict with 'id' and 'content'/'data' keys - + Examples: output('explore_0') # Read as raw text output('reviewer_0', format='json') # Read with metadata @@ -180,33 +181,35 @@ if "__omp_prelude_loaded__" not in globals(): raise RuntimeError("No session - output artifacts unavailable") artifacts_dir = session_file.rsplit(".", 1)[0] # Strip .jsonl extension if not Path(artifacts_dir).exists(): - _emit_status("output", error="Artifacts directory not found", path=artifacts_dir) + _emit_status( + "output", error="Artifacts directory not found", path=artifacts_dir + ) raise RuntimeError(f"No artifacts directory found: {artifacts_dir}") - + if not ids: _emit_status("output", error="No IDs provided") raise ValueError("At least one output ID is required") - + if query and (offset is not None or limit is not None): _emit_status("output", error="query cannot be combined with offset/limit") raise ValueError("query cannot be combined with offset/limit") - + results: list[dict] = [] not_found: list[str] = [] - + for output_id in ids: output_path = Path(artifacts_dir) / f"{output_id}.md" if not output_path.exists(): not_found.append(output_id) continue - + raw_content = output_path.read_text(encoding="utf-8") raw_lines = raw_content.splitlines() total_lines = len(raw_lines) - + selected_content = raw_content range_info: dict | None = None - + # Handle query if query: try: @@ -214,39 +217,60 @@ if "__omp_prelude_loaded__" not in globals(): except json.JSONDecodeError as e: _emit_status("output", id=output_id, error=f"Not valid JSON: {e}") raise ValueError(f"Output {output_id} is not valid JSON: {e}") - + # Apply jq-like query result_value = _apply_query(json_value, query) try: - selected_content = json.dumps(result_value, indent=2) if result_value is not None else "null" + selected_content = ( + json.dumps(result_value, indent=2) + if result_value is not None + else "null" + ) except (TypeError, ValueError): selected_content = str(result_value) - + # Handle offset/limit elif offset is not None or limit is not None: start_line = max(1, offset or 1) if start_line > total_lines: - _emit_status("output", id=output_id, error=f"Offset {start_line} beyond end ({total_lines} lines)") - raise ValueError(f"Offset {start_line} is beyond end of output ({total_lines} lines) for {output_id}") - - effective_limit = limit if limit is not None else total_lines - start_line + 1 + _emit_status( + "output", + id=output_id, + error=f"Offset {start_line} beyond end ({total_lines} lines)", + ) + raise ValueError( + f"Offset {start_line} is beyond end of output ({total_lines} lines) for {output_id}" + ) + + effective_limit = ( + limit if limit is not None else total_lines - start_line + 1 + ) end_line = min(total_lines, start_line + effective_limit - 1) selected_lines = raw_lines[start_line - 1 : end_line] selected_content = "\n".join(selected_lines) - range_info = {"start_line": start_line, "end_line": end_line, "total_lines": total_lines} - + range_info = { + "start_line": start_line, + "end_line": end_line, + "total_lines": total_lines, + } + # Strip ANSI codes if requested if format == "stripped": import re + selected_content = re.sub(r"\x1b\[[0-9;]*m", "", selected_content) - + # Build result if format == "json": result_data = { "id": output_id, "path": str(output_path), - "line_count": total_lines if not query else len(selected_content.splitlines()), - "char_count": len(raw_content) if not query else len(selected_content), + "line_count": total_lines + if not query + else len(selected_content.splitlines()), + "char_count": len(raw_content) + if not query + else len(selected_content), "content": selected_content, } if range_info: @@ -256,12 +280,10 @@ if "__omp_prelude_loaded__" not in globals(): results.append(result_data) else: results.append({"id": output_id, "content": selected_content}) - + # Handle not found if not_found: - available = sorted( - [f.stem for f in Path(artifacts_dir).glob("*.md")] - ) + available = sorted([f.stem for f in Path(artifacts_dir).glob("*.md")]) error_msg = f"Output not found: {', '.join(not_found)}" if available: error_msg += f"\n\nAvailable outputs: {', '.join(available[:20])}" @@ -269,7 +291,7 @@ if "__omp_prelude_loaded__" not in globals(): error_msg += f" (and {len(available) - 20} more)" _emit_status("output", not_found=not_found, available_count=len(available)) raise FileNotFoundError(error_msg) - + # Return format if len(ids) == 1: if format == "json": @@ -277,13 +299,13 @@ if "__omp_prelude_loaded__" not in globals(): return results[0] _emit_status("output", id=ids[0], chars=len(results[0]["content"])) return results[0]["content"] - + # Multiple IDs if format == "json": total_chars = sum(r["char_count"] for r in results) _emit_status("output", count=len(results), total_chars=total_chars) return results - + combined_output: list[dict] = [] for r in results: combined_output.append({"id": r["id"], "content": r["content"]}) @@ -295,13 +317,13 @@ if "__omp_prelude_loaded__" not in globals(): """Apply jq-like query to data. Supports .key, [index], and chaining.""" if not query: return data - + query = query.strip() if query.startswith("."): query = query[1:] if not query: return data - + # Parse query into tokens tokens = [] current_token = "" @@ -320,7 +342,7 @@ if "__omp_prelude_loaded__" not in globals(): j = i + 1 while j < len(query) and query[j] != "]": j += 1 - bracket_content = query[i+1:j] + bracket_content = query[i + 1 : j] if bracket_content.startswith('"') and bracket_content.endswith('"'): tokens.append(("key", bracket_content[1:-1])) else: @@ -331,7 +353,7 @@ if "__omp_prelude_loaded__" not in globals(): i += 1 if current_token: tokens.append(("key", current_token)) - + # Apply tokens current = data for token_type, value in tokens: @@ -343,9 +365,8 @@ if "__omp_prelude_loaded__" not in globals(): if not isinstance(current, dict) or value not in current: return None current = current[value] - - return current + return current def _tool_proxy_from_env() -> tuple[str, str, str]: base = os.environ.get("PI_TOOL_BRIDGE_URL") @@ -358,9 +379,14 @@ if "__omp_prelude_loaded__" not in globals(): def _bridge_call(name: str, args: dict): """POST one request to the host tool bridge and return its `value`.""" import urllib.request, urllib.error + base, token, session = _tool_proxy_from_env() _run_id_getter = globals().get("__omp_current_run_id__") - _run_id = _run_id_getter() if callable(_run_id_getter) else globals().get("__omp_run_id__") + _run_id = ( + _run_id_getter() + if callable(_run_id_getter) + else globals().get("__omp_run_id__") + ) payload = json.dumps( {"session": session, "run": _run_id, "name": name, "args": args} ).encode("utf-8") @@ -429,7 +455,11 @@ if "__omp_prelude_loaded__" not in globals(): def __repr__(self) -> str: session = os.environ.get("PI_TOOL_BRIDGE_SESSION") - return f"" if session else "" + return ( + f"" + if session + else "" + ) tool = _ToolProxy() @@ -450,7 +480,18 @@ if "__omp_prelude_loaded__" not in globals(): text = res.get("text") if isinstance(res, dict) else res return json.loads(text) if schema is not None else text - def agent(prompt, *, agent="task", model=None, label=None, schema=None, isolated=None, apply=None, merge=None, handle=False): + def agent( + prompt, + *, + agent="task", + model=None, + label=None, + schema=None, + isolated=None, + apply=None, + merge=None, + handle=False, + ): """Run a subagent and return its final output. `agent` selects the subagent definition (default "task"). Pass @@ -513,7 +554,13 @@ if "__omp_prelude_loaded__" not in globals(): return parsed details = res.get("details") if isinstance(res, dict) else None if not isinstance(details, dict) or details.get("id") is None: - return {"text": text, "output": text, "handle": None, "id": None, "agent": None} + return { + "text": text, + "output": text, + "handle": None, + "id": None, + "agent": None, + } node = { "text": text, "output": text, @@ -559,6 +606,7 @@ if "__omp_prelude_loaded__" not in globals(): pool width tracks ``task.maxConcurrency`` (0 = run every item at once). """ import concurrent.futures, contextvars + items = list(items) if not items: return [] diff --git a/packages/coding-agent/src/eval/py/runner.py b/packages/coding-agent/src/eval/py/runner.py index 9bc1a7ac5..e9f116728 100644 --- a/packages/coding-agent/src/eval/py/runner.py +++ b/packages/coding-agent/src/eval/py/runner.py @@ -191,10 +191,14 @@ class _RunnerState: self.capture_rid: str | None = None -_CURRENT_RID: contextvars.ContextVar[str | None] = contextvars.ContextVar("omp_current_rid", default=None) -_CURRENT_DISPLAYED_MATPLOTLIB_FIGURE_IDS: contextvars.ContextVar[set[int] | None] = contextvars.ContextVar( - "omp_displayed_matplotlib_figure_ids", - default=None, +_CURRENT_RID: contextvars.ContextVar[str | None] = contextvars.ContextVar( + "omp_current_rid", default=None +) +_CURRENT_DISPLAYED_MATPLOTLIB_FIGURE_IDS: contextvars.ContextVar[set[int] | None] = ( + contextvars.ContextVar( + "omp_displayed_matplotlib_figure_ids", + default=None, + ) ) @@ -233,7 +237,9 @@ def _drain_captured_stdout() -> None: def _start_capture_drain() -> None: if _CAPTURE_READ_FD is None: return - thread = threading.Thread(target=_drain_captured_stdout, name="omp-fd1-capture", daemon=True) + thread = threading.Thread( + target=_drain_captured_stdout, name="omp-fd1-capture", daemon=True + ) thread.start() @@ -242,7 +248,9 @@ def _start_capture_drain() -> None: # --------------------------------------------------------------------------- -_MAGIC_LINE_RE = re.compile(r"^(?P[ \t]*)(?P[A-Za-z_][A-Za-z_0-9]*)(?:[ \t]+(?P.*))?$") +_MAGIC_LINE_RE = re.compile( + r"^(?P[ \t]*)(?P[A-Za-z_][A-Za-z_0-9]*)(?:[ \t]+(?P.*))?$" +) _ASSIGN_LINE_RE = re.compile( r"^(?P[ \t]*)(?P[A-Za-z_][A-Za-z_0-9.\[\], ]*?)\s*=\s*(?P.+)$" ) @@ -337,7 +345,9 @@ def transform_cell(source: str) -> str: rhs = m.group("rhs").strip() if rhs.startswith("!"): cmd = rhs[1:].strip() - out.append(f"{m.group('indent')}{m.group('lhs').rstrip()} = __omp_shell({_quote_arg(cmd)})") + out.append( + f"{m.group('indent')}{m.group('lhs').rstrip()} = __omp_shell({_quote_arg(cmd)})" + ) i += 1 continue if rhs.startswith("%") and not rhs.startswith("%%"): @@ -383,7 +393,9 @@ def line_magic(name: str) -> Callable[[Callable[[str], Any]], Callable[[str], An return decorator -def cell_magic(name: str) -> Callable[[Callable[[str, str], Any]], Callable[[str, str], Any]]: +def cell_magic( + name: str, +) -> Callable[[Callable[[str, str], Any]], Callable[[str, str], Any]]: def decorator(fn: Callable[[str, str], Any]) -> Callable[[str, str], Any]: _CELL_MAGICS[name] = fn return fn @@ -398,6 +410,7 @@ def _emit_status(op: str, **data: Any) -> None: return _emit({"type": "display", "id": rid, "bundle": bundle}) + _SHELL_READ_CHUNK_BYTES = 8192 _SHELL_OUTPUT_MAX_BYTES = 1024 * 1024 _SHELL_OUTPUT_MAX_LINES = 3000 @@ -458,12 +471,16 @@ class _ShellOutputLimiter: return limited = _take_prefix_by_lines(text, self._remaining_lines) truncated = limited != text - byte_limited = _take_prefix_by_encoded_bytes(limited, self._remaining_bytes, self._encoding) + byte_limited = _take_prefix_by_encoded_bytes( + limited, self._remaining_bytes, self._encoding + ) truncated = truncated or byte_limited != limited if byte_limited: sys.stdout.write(byte_limited) sys.stdout.flush() - self._remaining_bytes -= len(byte_limited.encode(self._encoding, errors="strict")) + self._remaining_bytes -= len( + byte_limited.encode(self._encoding, errors="strict") + ) self._remaining_lines -= byte_limited.count("\n") self._at_line_start = byte_limited.endswith("\n") if truncated: @@ -478,7 +495,9 @@ class _ShellOutputLimiter: self._truncated = True -def _stream_process_output(proc: subprocess.Popen, on_text: Callable[[str], None] | None = None) -> None: +def _stream_process_output( + proc: subprocess.Popen, on_text: Callable[[str], None] | None = None +) -> None: assert proc.stdout is not None encoding = _process_output_encoding() decoder = _process_output_decoder(encoding) @@ -514,7 +533,9 @@ class _BoundedTextCapture: if self._remaining_bytes <= 0 or self._remaining_lines <= 0: return line_limited = _take_prefix_by_lines(text, self._remaining_lines) - part = _take_prefix_by_encoded_bytes(line_limited, self._remaining_bytes, self._encoding) + part = _take_prefix_by_encoded_bytes( + line_limited, self._remaining_bytes, self._encoding + ) if not part: return self._parts.append(part) @@ -583,7 +604,9 @@ def _magic_pip(args: str) -> None: head = mod_name.split(".", 1)[0].lower() if head in prefixes: sys.modules.pop(mod_name, None) - _emit_status("pip", args=args, installed=installed_packages, exit_code=proc.returncode) + _emit_status( + "pip", args=args, installed=installed_packages, exit_code=proc.returncode + ) @line_magic("cd") @@ -658,7 +681,9 @@ def _magic_who(_args: str) -> list[str]: names = sorted( name for name, value in _STATE.user_ns.items() - if not name.startswith("_") and not callable(value) or hasattr(value, "__class__") + if not name.startswith("_") + and not callable(value) + or hasattr(value, "__class__") ) return [n for n in names if not n.startswith("__")] @@ -677,7 +702,9 @@ def _magic_whos(_args: str) -> list[tuple[str, str]]: @line_magic("reset") def _magic_reset(_args: str) -> None: _STATE.user_ns.clear() - _STATE.user_ns.update({"__name__": "__main__", "__doc__": None, "__builtins__": builtins}) + _STATE.user_ns.update( + {"__name__": "__main__", "__doc__": None, "__builtins__": builtins} + ) _install_builtins(_STATE.user_ns) _emit_status("reset") @@ -686,7 +713,9 @@ def _magic_reset(_args: str) -> None: def _magic_load(args: str) -> None: path = Path(os.path.expanduser(args.strip())) source = path.read_text(encoding="utf-8") - _emit({"type": "display", "id": _CURRENT_RID.get(), "bundle": {"text/plain": source}}) + _emit( + {"type": "display", "id": _CURRENT_RID.get(), "bundle": {"text/plain": source}} + ) _exec_source(source, _STATE.user_ns) @@ -712,6 +741,7 @@ def _magic_run(args: str) -> None: def _magic_cell_bash(args: str, body: str) -> int: return _run_shell_body(body, shell_arg="/bin/bash") + @cell_magic("capture") def _magic_cell_capture(args: str, body: str) -> str: """Capture stdout/stderr of body; bind to ``args`` (a name) if provided.""" @@ -797,7 +827,9 @@ def __omp_shell(cmd: str) -> _ShellResult: stdout=subprocess.PIPE, stderr=subprocess.STDOUT, ) - capture = _BoundedTextCapture(_SHELL_RESULT_CAPTURE_BYTES, _SHELL_OUTPUT_MAX_LINES, _process_output_encoding()) + capture = _BoundedTextCapture( + _SHELL_RESULT_CAPTURE_BYTES, _SHELL_OUTPUT_MAX_LINES, _process_output_encoding() + ) _stream_process_output(proc, capture.add) proc.wait() lines = [line for line in capture.text().splitlines()] @@ -827,7 +859,9 @@ def _is_matplotlib_figure(value: Any) -> bool: return True value_type = type(value) - return value_type.__module__ == "matplotlib.figure" and value_type.__name__ == "Figure" + return ( + value_type.__module__ == "matplotlib.figure" and value_type.__name__ == "Figure" + ) def _matplotlib_figure_png(value: Any) -> str | None: @@ -869,7 +903,6 @@ def _mime_bundle(value: Any) -> dict: if matplotlib_png is not None: bundle["image/png"] = matplotlib_png - mimebundle = getattr(value, "_repr_mimebundle_", None) if callable(mimebundle): try: @@ -990,7 +1023,9 @@ def _await_sync(coro) -> Any: except RuntimeError: running_loop = None if running_loop is not None and running_loop.is_running(): - raise RuntimeError("top-level await is not supported from synchronous magic execution") + raise RuntimeError( + "top-level await is not supported from synchronous magic execution" + ) return asyncio.run(coro) @@ -1005,7 +1040,6 @@ def _run_compiled_sync(code, ns: dict, *, want_value: bool) -> Any: return None - async def _run_compiled_async(code, ns: dict, *, want_value: bool) -> Any: """Execute a code object in the persistent event loop. @@ -1127,6 +1161,7 @@ def _apply_request_runtime(req: dict) -> None: elif value is None: os.environ.pop(key, None) + def _start_parent_watchdog() -> None: """Self-terminate when the host process dies. @@ -1182,23 +1217,27 @@ async def _handle_request_async(req: dict) -> None: transformed = transform_cell(req.get("code", "")) except SyntaxError as exc: _emit_error(rid, exc) - _emit({ - "type": "done", - "id": rid, - "status": "error", - "executionCount": execution_count, - "cancelled": False, - }) + _emit( + { + "type": "done", + "id": rid, + "status": "error", + "executionCount": execution_count, + "cancelled": False, + } + ) return except BaseException as exc: # noqa: BLE001 - runtime setup errors must settle the request _emit_error(rid, exc) - _emit({ - "type": "done", - "id": rid, - "status": "error", - "executionCount": execution_count, - "cancelled": False, - }) + _emit( + { + "type": "done", + "id": rid, + "status": "error", + "executionCount": execution_count, + "cancelled": False, + } + ) return _begin_exec_sigint() @@ -1222,13 +1261,15 @@ async def _handle_request_async(req: dict) -> None: pass _flush_stream_proxies(rid) - _emit({ - "type": "done", - "id": rid, - "status": status, - "executionCount": execution_count, - "cancelled": cancelled, - }) + _emit( + { + "type": "done", + "id": rid, + "status": status, + "executionCount": execution_count, + "cancelled": cancelled, + } + ) finally: if _STATE.capture_rid == rid: _STATE.capture_rid = None @@ -1239,13 +1280,15 @@ async def _handle_request_async(req: dict) -> None: def _emit_error(rid: str, exc: BaseException) -> None: tb_lines = traceback.format_exception(type(exc), exc, exc.__traceback__) - _emit({ - "type": "error", - "id": rid, - "ename": type(exc).__name__, - "evalue": str(exc), - "traceback": [line.rstrip("\n") for line in tb_lines], - }) + _emit( + { + "type": "error", + "id": rid, + "ename": type(exc).__name__, + "evalue": str(exc), + "traceback": [line.rstrip("\n") for line in tb_lines], + } + ) # --------------------------------------------------------------------------- @@ -1261,13 +1304,15 @@ def _read_stdin(loop: asyncio.AbstractEventLoop, queue: asyncio.Queue, stdin) -> try: req = json.loads(line) except json.JSONDecodeError as exc: - _emit({ - "type": "error", - "id": "", - "ename": "ProtocolError", - "evalue": f"Invalid JSON request: {exc}", - "traceback": [], - }) + _emit( + { + "type": "error", + "id": "", + "ename": "ProtocolError", + "evalue": f"Invalid JSON request: {exc}", + "traceback": [], + } + ) continue loop.call_soon_threadsafe(queue.put_nowait, req) loop.call_soon_threadsafe(queue.put_nowait, {"type": "exit"}) @@ -1287,10 +1332,16 @@ async def _main_async() -> None: loop = asyncio.get_running_loop() _STATE.loop = loop queue: asyncio.Queue = asyncio.Queue() - reader = threading.Thread(target=_read_stdin, args=(loop, queue, stdin), name="omp-stdin-reader", daemon=True) + reader = threading.Thread( + target=_read_stdin, + args=(loop, queue, stdin), + name="omp-stdin-reader", + daemon=True, + ) reader.start() tasks: set[asyncio.Task] = set() + def _task_done(task: asyncio.Task) -> None: tasks.discard(task) try: @@ -1299,6 +1350,7 @@ async def _main_async() -> None: return if exc is not None: _emit_error("", exc) + try: while True: req = await queue.get() diff --git a/packages/metaharness/agent/omp_local.py b/packages/metaharness/agent/omp_local.py index 46c4fb9d0..edffd4678 100644 --- a/packages/metaharness/agent/omp_local.py +++ b/packages/metaharness/agent/omp_local.py @@ -153,6 +153,7 @@ def _env(name: str, default: str = "") -> str: def _truthy(value: str) -> bool: return value.strip().lower() in ("1", "true", "yes", "on") + def _loads(line: str) -> dict | None: line = line.strip() if not line.startswith("{"): @@ -201,11 +202,15 @@ class OmpLocal(BaseInstalledAgent): self._tarball = _env("OMP_BENCH_TARBALL") self._pkg_version = _env("OMP_BENCH_VERSION", "latest") self._models_yaml_path = _env("OMP_BENCH_MODELS_YAML") - self._gateway_url = _env("OMP_BENCH_GATEWAY_URL", "http://host.docker.internal:4000") + self._gateway_url = _env( + "OMP_BENCH_GATEWAY_URL", "http://host.docker.internal:4000" + ) self._gateway_token = _env("OMP_BENCH_GATEWAY_TOKEN", "no-auth-dummy") self._gateway_providers = [ p.strip() - for p in _env("OMP_BENCH_GATEWAY_PROVIDERS", "anthropic,openai-codex").split(",") + for p in _env( + "OMP_BENCH_GATEWAY_PROVIDERS", "anthropic,openai-codex" + ).split(",") if p.strip() ] self._thinking = _env("OMP_BENCH_THINKING") @@ -248,7 +253,9 @@ class OmpLocal(BaseInstalledAgent): def get_version_command(self) -> str | None: if self._binary: return f"{shlex.quote(self._cli)} --version" - return self._wrap(f"{shlex.quote(self._bun)} {shlex.quote(self._cli)} --version") + return self._wrap( + f"{shlex.quote(self._bun)} {shlex.quote(self._cli)} --version" + ) @override def parse_version(self, stdout: str) -> str: @@ -264,7 +271,7 @@ class OmpLocal(BaseInstalledAgent): """ bun_dir = os.path.dirname(self._bun) return ( - f'export BUN_INSTALL={shlex.quote(self._home + "/.bun")}; ' + f"export BUN_INSTALL={shlex.quote(self._home + '/.bun')}; " f'export PATH="{bun_dir}:$PATH"; ' f"{command}" ) @@ -272,7 +279,9 @@ class OmpLocal(BaseInstalledAgent): @override async def install(self, environment: BaseEnvironment) -> None: # Resolve the agent user's HOME first (root vs non-root tasks differ). - home = (await self.exec_as_agent(environment, command='printf %s "$HOME"')).stdout + home = ( + await self.exec_as_agent(environment, command='printf %s "$HOME"') + ).stdout self._home = (home or "/root").strip() or "/root" if self._binary: @@ -304,7 +313,7 @@ class OmpLocal(BaseInstalledAgent): "set -e; " f"export BUN_INSTALL={shlex.quote(self._home + '/.bun')}; " f'curl -fsSL https://bun.sh/install | bash -s "bun-v{self._bun_version}"; ' - f'{shlex.quote(self._home + "/.bun/bin/bun")} --version' + f"{shlex.quote(self._home + '/.bun/bin/bun')} --version" ), ) self._bun = f"{self._home}/.bun/bin/bun" @@ -326,8 +335,15 @@ class OmpLocal(BaseInstalledAgent): `node_modules` with a linux tree, and mounts a linux `bun` binary — so setup needs zero outbound network and no rebuild for TS changes. """ - arch = (await self.exec_as_agent(environment, command="uname -m")).stdout.strip() - norm = {"aarch64": "arm64", "arm64": "arm64", "x86_64": "x64", "amd64": "x64"}.get(arch) + arch = ( + await self.exec_as_agent(environment, command="uname -m") + ).stdout.strip() + norm = { + "aarch64": "arm64", + "arm64": "arm64", + "x86_64": "x64", + "amd64": "x64", + }.get(arch) if self._source_arch and norm != self._source_arch: raise RuntimeError( f"source mode: container arch {arch!r} != mounted deps tree arch " @@ -351,7 +367,9 @@ class OmpLocal(BaseInstalledAgent): async def _install_local(self, environment: BaseEnvironment) -> str: if not self._tarball: - raise RuntimeError("OMP_BENCH_INSTALL=local requires OMP_BENCH_TARBALL (host tarball path)") + raise RuntimeError( + "OMP_BENCH_INSTALL=local requires OMP_BENCH_TARBALL (host tarball path)" + ) await environment.upload_file(self._tarball, _TARBALL_DST) app = f"{self._home}/.omp-bench/app" await self.exec_as_agent( @@ -364,7 +382,7 @@ class OmpLocal(BaseInstalledAgent): # Bundle inlines workspace TS; only externalized deps are needed. # Skip heavy optionals (transformers/sherpa) but add the native addon. "bun install --production --omit=optional; " - 'arch=$(uname -m); ' + "arch=$(uname -m); " 'case "$arch" in aarch64|arm64) na=arm64 ;; x86_64|amd64) na=x64 ;; ' '*) echo "unsupported arch $arch" >&2; exit 4 ;; esac; ' # Native leaf MUST match the bundle version exactly (loader/API skew @@ -379,7 +397,9 @@ class OmpLocal(BaseInstalledAgent): async def _install_binary(self, environment: BaseEnvironment) -> str: """Probe container arch, upload only the matching self-contained omp binary.""" - arch = (await self.exec_as_agent(environment, command="uname -m")).stdout.strip() + arch = ( + await self.exec_as_agent(environment, command="uname -m") + ).stdout.strip() if arch in ("aarch64", "arm64"): hostbin = self._binary_arm64 elif arch in ("x86_64", "amd64"): @@ -387,11 +407,15 @@ class OmpLocal(BaseInstalledAgent): else: raise RuntimeError(f"binary mode: unsupported container arch {arch!r}") if not hostbin: - raise RuntimeError(f"binary mode: no omp binary provided for container arch {arch}") + raise RuntimeError( + f"binary mode: no omp binary provided for container arch {arch}" + ) app_dir = f"{self._home}/.omp-bench" dst = f"{app_dir}/omp" staging = "/tmp/omp-bin" - await self.exec_as_agent(environment, command=f"mkdir -p {shlex.quote(app_dir)}") + await self.exec_as_agent( + environment, command=f"mkdir -p {shlex.quote(app_dir)}" + ) await environment.upload_file(hostbin, staging) await self.exec_as_agent( environment, @@ -422,7 +446,9 @@ class OmpLocal(BaseInstalledAgent): else: content = self._generate_models_yaml() staged = _MODELS_DST - heredoc = f"cat > {_MODELS_DST} <<'OMP_MODELS_EOF'\n{content}\nOMP_MODELS_EOF" + heredoc = ( + f"cat > {_MODELS_DST} <<'OMP_MODELS_EOF'\n{content}\nOMP_MODELS_EOF" + ) await self.exec_as_agent(environment, command=heredoc) await self.exec_as_agent( environment, @@ -433,7 +459,10 @@ class OmpLocal(BaseInstalledAgent): ) def _generate_models_yaml(self) -> str: - lines = ["# Generated by metaharness runner — routes auth via host gateway.", "providers:"] + lines = [ + "# Generated by metaharness runner — routes auth via host gateway.", + "providers:", + ] for provider in self._gateway_providers: lines += [ f" {provider}:", @@ -513,7 +542,9 @@ class OmpLocal(BaseInstalledAgent): context: AgentContext, ) -> None: if not self.model_name or "/" not in self.model_name: - raise ValueError("model must be 'provider/model' (e.g. anthropic/claude-sonnet-4-6)") + raise ValueError( + "model must be 'provider/model' (e.g. anthropic/claude-sonnet-4-6)" + ) provider, model = self.model_name.split("/", 1) if self._binary: @@ -547,7 +578,11 @@ class OmpLocal(BaseInstalledAgent): if not self._gateway_on: run_env.update(self._collect_provider_keys(provider)) run_env.update(self._forward_env) - await self.exec_as_agent(environment, command=run if self._binary else self._wrap(run), env=run_env or None) + await self.exec_as_agent( + environment, + command=run if self._binary else self._wrap(run), + env=run_env or None, + ) @override def populate_context_post_run(self, context: AgentContext) -> None: diff --git a/packages/metaharness/src/adapters/snapcompact.py b/packages/metaharness/src/adapters/snapcompact.py index b45daf69b..c530fcb29 100644 --- a/packages/metaharness/src/adapters/snapcompact.py +++ b/packages/metaharness/src/adapters/snapcompact.py @@ -38,6 +38,7 @@ from pathlib import Path HERE = Path(__file__).resolve().parent RESEARCH = HERE.parents[2] / "snapcompact" / "research" + def find_agent_prompts() -> Path: for parent in HERE.parents: for candidate in ( @@ -110,7 +111,9 @@ def agent_prompt(name: str) -> str: return re.sub(r"\{\{#if .*?\{\{/if\}\}\n?", "", text, flags=re.DOTALL) -def cached_complete(api_key: str, model: str, messages: list[dict], fresh: bool, **kw) -> tuple[str, dict]: +def cached_complete( + api_key: str, model: str, messages: list[dict], fresh: bool, **kw +) -> tuple[str, dict]: """complete() with response caching keyed on the full request payload. Truncated responses (stop_reason == max_tokens) are never cached and never @@ -135,8 +138,15 @@ def parse_condition(name: str) -> dict: return {"name": name, "kind": name} m = re.fullmatch(r"img-([a-z0-9]+)-([a-z-]+)", name) if not m or m.group(1) not in FONTS or m.group(2) not in VARIANTS: - raise SystemExit(f"bad condition {name!r}; expected text|compact|handoff|img--") - return {"name": name, "kind": "image", "font": FONTS[m.group(1)], "variant": m.group(2)} + raise SystemExit( + f"bad condition {name!r}; expected text|compact|handoff|img--" + ) + return { + "name": name, + "kind": "image", + "font": FONTS[m.group(1)], + "variant": m.group(2), + } def run_chunk(cond: dict, start: int, end: int, ctx_args: dict) -> list[dict]: @@ -148,7 +158,9 @@ def run_chunk(cond: dict, start: int, end: int, ctx_args: dict) -> list[dict]: ctx_args["offsets"], ctx_args["api_key"], ) - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] chunk_text = flow[start:end] @@ -157,21 +169,41 @@ def run_chunk(cond: dict, start: int, end: int, ctx_args: dict) -> list[dict]: png = cols = rows = None context = chunk_text if cond["kind"] == "image": - salt = ("dimv2",) if cond["variant"] == "dim" else () # cache-bust pre-fix sticky-fg dim renders - png = CACHE / f"img-{cond['font'].name}-{cond['variant']}-{sha8(chunk_text, str(args.size), *salt)}.png" + salt = ( + ("dimv2",) if cond["variant"] == "dim" else () + ) # cache-bust pre-fix sticky-fg dim renders + png = ( + CACHE + / f"img-{cond['font'].name}-{cond['variant']}-{sha8(chunk_text, str(args.size), *salt)}.png" + ) if not png.exists(): - render(chunk_text, cond["font"], CACHE, args.size, cond["variant"]).save(png) + render(chunk_text, cond["font"], CACHE, args.size, cond["variant"]).save( + png + ) cols, rows, _ = capacity(cond["font"], args.size) elif cond["kind"] in ("compact", "handoff"): - prompt_file = {"compact": "compaction-summary.md", "handoff": "handoff-document.md"}[cond["kind"]] + prompt_file = { + "compact": "compaction-summary.md", + "handoff": "handoff-document.md", + }[cond["kind"]] gen_messages = [ - {"role": "user", "content": load_prompt("session-frame.md").format(context=chunk_text)}, - {"role": "assistant", "content": "Noted. I have read the passages and will keep them in mind."}, + { + "role": "user", + "content": load_prompt("session-frame.md").format(context=chunk_text), + }, + { + "role": "assistant", + "content": "Noted. I have read the passages and will keep them in mind.", + }, {"role": "user", "content": agent_prompt(prompt_file)}, ] context, gen_usage = cached_complete( - api_key, args.model, gen_messages, args.fresh, - system=agent_prompt("summarization-system.md"), max_tokens=4096, + api_key, + args.model, + gen_messages, + args.fresh, + system=agent_prompt("summarization-system.md"), + max_tokens=4096, ) usage_rows.append(("summarize", gen_usage)) @@ -183,16 +215,30 @@ def run_chunk(cond: dict, start: int, end: int, ctx_args: dict) -> list[dict]: q_block = "\n".join(f"{i + 1}. {q['q']}" for i, q in enumerate(batch)) if cond["kind"] == "image": carrier = image_block(png) - preamble = {"type": "text", "text": load_prompt("qa-image.md").format(cols=cols, rows=rows)} + preamble = { + "type": "text", + "text": load_prompt("qa-image.md").format(cols=cols, rows=rows), + } else: - carrier = {"type": "text", "text": load_prompt("qa-text.md").format(context=context)} + carrier = { + "type": "text", + "text": load_prompt("qa-text.md").format(context=context), + } preamble = None if use_cache: carrier["cache_control"] = {"type": "ephemeral"} - content = ([preamble] if preamble else []) + [carrier, {"type": "text", "text": q_block}] + content = ([preamble] if preamble else []) + [ + carrier, + {"type": "text", "text": q_block}, + ] messages = [{"role": "user", "content": content}] text, usage = cached_complete( - api_key, args.model, messages, args.fresh, max_tokens=args.max_tokens, effort=args.effort + api_key, + args.model, + messages, + args.fresh, + max_tokens=args.max_tokens, + effort=args.effort, ) usage_rows.append(("qa", usage)) answers.extend(squad.parse_numbered(text, len(batch))) @@ -225,7 +271,9 @@ def run_chunk(cond: dict, start: int, end: int, ctx_args: dict) -> list[dict]: return records -def aggregate(name: str, records: list[dict], price_in: float, price_out: float) -> dict: +def aggregate( + name: str, records: list[dict], price_in: float, price_out: float +) -> dict: n = len(records) f1s = [r["f1"] for r in records] mean_f1 = sum(f1s) / n @@ -251,7 +299,8 @@ def aggregate(name: str, records: list[dict], price_in: float, price_out: float) "cache_w": cache_w, "cache_r": cache_r, # Anthropic pricing: cache write 1.25x input, cache read 0.1x input (5m TTL). - "cost_usd": (tok_in + 1.25 * cache_w + 0.1 * cache_r) / 1e6 * price_in + tok_out / 1e6 * price_out, + "cost_usd": (tok_in + 1.25 * cache_w + 0.1 * cache_r) / 1e6 * price_in + + tok_out / 1e6 * price_out, "f1_by_quartile": quart, } @@ -261,23 +310,59 @@ def main() -> None: ap.add_argument("--model", default="claude-fable-5") ap.add_argument("--conditions", default=DEFAULT_CONDITIONS) ap.add_argument("--qpc", type=int, default=30, help="questions sampled per chunk") - ap.add_argument("--qpb", type=int, default=0, help="questions per API call (batches the chunk); 0 = all at once") - ap.add_argument("--cache", choices=["auto", "on", "off"], default="auto", - help="prompt-cache the carrier block; auto = on when --qpb is set") - ap.add_argument("--max-tokens", type=int, default=8192, help="output budget per QA call (incl. thinking)") - ap.add_argument("--effort", choices=["low", "medium", "high", "xhigh", "max"], default=None, - help="adaptive-thinking effort for QA calls; default = provider default") + ap.add_argument( + "--qpb", + type=int, + default=0, + help="questions per API call (batches the chunk); 0 = all at once", + ) + ap.add_argument( + "--cache", + choices=["auto", "on", "off"], + default="auto", + help="prompt-cache the carrier block; auto = on when --qpb is set", + ) + ap.add_argument( + "--max-tokens", + type=int, + default=8192, + help="output budget per QA call (incl. thinking)", + ) + ap.add_argument( + "--effort", + choices=["low", "medium", "high", "xhigh", "max"], + default=None, + help="adaptive-thinking effort for QA calls; default = provider default", + ) ap.add_argument("--seed", type=int, default=42) ap.add_argument("--size", type=int, default=1568) ap.add_argument("--workers", type=int, default=4) - ap.add_argument("--limit-chars", type=int, default=0, help="cap corpus size; 0 = full dev set") - ap.add_argument("--limit-paras", type=int, default=0, help="cap corpus to first N passages; 0 = all") + ap.add_argument( + "--limit-chars", type=int, default=0, help="cap corpus size; 0 = full dev set" + ) + ap.add_argument( + "--limit-paras", + type=int, + default=0, + help="cap corpus to first N passages; 0 = all", + ) ap.add_argument("--fresh", action="store_true", help="ignore cached responses") - ap.add_argument("--report", action="store_true", help="aggregate cached records only; no API calls") - ap.add_argument("--price-in", type=float, default=10.0, help="$ per 1M input tokens") - ap.add_argument("--price-out", type=float, default=50.0, help="$ per 1M output tokens") + ap.add_argument( + "--report", + action="store_true", + help="aggregate cached records only; no API calls", + ) + ap.add_argument( + "--price-in", type=float, default=10.0, help="$ per 1M input tokens" + ) + ap.add_argument( + "--price-out", type=float, default=50.0, help="$ per 1M output tokens" + ) ap.add_argument("--env", default="~/.env") - ap.add_argument("--output-dir", help="write records.jsonl and summary.json directly to this directory") + ap.add_argument( + "--output-dir", + help="write records.jsonl and summary.json directly to this directory", + ) args = ap.parse_args() CACHE.mkdir(exist_ok=True) @@ -301,26 +386,43 @@ def main() -> None: if args.limit_paras: paras = paras[: args.limit_paras] flow, offsets = squad.build_flow(paras, args.limit_chars or None) - conditions = [parse_condition(c.strip()) for c in args.conditions.split(",") if c.strip()] + conditions = [ + parse_condition(c.strip()) for c in args.conditions.split(",") if c.strip() + ] tasks: list[tuple[dict, int, int]] = [] for cond in conditions: - budget = capacity(cond["font"], args.size)[2] if cond["kind"] == "image" else TEXT_CHUNK + budget = ( + capacity(cond["font"], args.size)[2] + if cond["kind"] == "image" + else TEXT_CHUNK + ) for start in range(0, len(flow), budget): tasks.append((cond, start, min(start + budget, len(flow)))) - calls = len(tasks) + sum(1 for c, *_ in tasks if c["kind"] in ("compact", "handoff")) + calls = len(tasks) + sum( + 1 for c, *_ in tasks if c["kind"] in ("compact", "handoff") + ) print( f"corpus={len(flow):,} chars ({len(offsets):,} passages), {len(conditions)} conditions, " f"{len(tasks)} chunks, <= {calls} API calls, qpc={args.qpc}, model={args.model}" ) api_key = "" if args.report else load_api_key(args.env) - ctx_args = {"args": args, "flow": flow, "paras": paras, "offsets": offsets, "api_key": api_key} + ctx_args = { + "args": args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "api_key": api_key, + } records: list[dict] = [] done = 0 with (run_dir / "records.jsonl").open("w") as records_file: with ThreadPoolExecutor(args.workers) as pool: - futures = [pool.submit(run_chunk, cond, start, end, ctx_args) for cond, start, end in tasks] + futures = [ + pool.submit(run_chunk, cond, start, end, ctx_args) + for cond, start, end in tasks + ] for fut in futures: chunk_records = fut.result() records.extend(chunk_records) @@ -332,12 +434,19 @@ def main() -> None: print(f" {done}/{len(tasks)} chunks", flush=True) rows = [ - aggregate(cond["name"], [r for r in records if r["cond"] == cond["name"]], args.price_in, args.price_out) + aggregate( + cond["name"], + [r for r in records if r["cond"] == cond["name"]], + args.price_in, + args.price_out, + ) for cond in conditions if any(r["cond"] == cond["name"] for r in records) ] rows.sort(key=lambda r: -r["f1"]) - (run_dir / "summary.json").write_text(json.dumps({"args": vars(args), "rows": rows}, indent=1)) + (run_dir / "summary.json").write_text( + json.dumps({"args": vars(args), "rows": rows}, indent=1) + ) hdr = ( f"{'condition':<15}{'n':>6}{'EM':>7}{'F1':>7}{'±se':>6}{'abst':>6}" @@ -351,7 +460,9 @@ def main() -> None: ) print(f"\n{'condition':<15} F1 by position quartile (Q1..Q4)") for r in rows: - cells = " ".join(" - " if q is None else f"{q:.3f}" for q in r["f1_by_quartile"]) + cells = " ".join( + " - " if q is None else f"{q:.3f}" for q in r["f1_by_quartile"] + ) print(f"{r['name']:<15} {cells}") print(f"\nresults -> {run_dir}/") diff --git a/packages/snapcompact/research/anthropic_api.py b/packages/snapcompact/research/anthropic_api.py index d6a94fa6e..cbe72a1de 100644 --- a/packages/snapcompact/research/anthropic_api.py +++ b/packages/snapcompact/research/anthropic_api.py @@ -27,7 +27,10 @@ def load_api_key(env_path: str = "~/.env") -> str: def image_block(png_path: Path) -> dict: data = base64.b64encode(png_path.read_bytes()).decode() - return {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": data}} + return { + "type": "image", + "source": {"type": "base64", "media_type": "image/png", "data": data}, + } def complete( @@ -61,7 +64,9 @@ def complete( try: with urllib.request.urlopen(req, timeout=300) as resp: out = json.load(resp) - text = "".join(b.get("text", "") for b in out["content"] if b.get("type") == "text") + text = "".join( + b.get("text", "") for b in out["content"] if b.get("type") == "text" + ) return text, out.get("usage", {}), out.get("stop_reason", "") except urllib.error.HTTPError as err: detail = err.read().decode(errors="replace")[:500] diff --git a/packages/snapcompact/research/bdf.py b/packages/snapcompact/research/bdf.py index 15003636d..e9118ee8f 100644 --- a/packages/snapcompact/research/bdf.py +++ b/packages/snapcompact/research/bdf.py @@ -9,7 +9,9 @@ from PIL import Image XORG_RAW = "https://gitlab.freedesktop.org/xorg/font/misc-misc/-/raw/master/{name}.bdf" TOM_THUMB = "https://robey.lag.net/downloads/tom-thumb.bdf" -UNSCII_HEX = "https://raw.githubusercontent.com/viznut/unscii/master/fontfiles/{name}.hex" +UNSCII_HEX = ( + "https://raw.githubusercontent.com/viznut/unscii/master/fontfiles/{name}.hex" +) @dataclass(frozen=True) @@ -21,7 +23,9 @@ class FontCfg: adv: int # x advance per character cell, px pitch: int # y advance per row, px ascent: int | None = None # override; default from FONT_ASCENT - native: tuple[int, int] | None = None # rasterize at this cell size, then resize (stretch) to adv x pitch + native: tuple[int, int] | None = ( + None # rasterize at this cell size, then resize (stretch) to adv x pitch + ) repeat: int = 1 # render each text line this many times (copy 0 plain, later copies bg-highlighted) @@ -32,7 +36,11 @@ def ensure_font(cfg: FontCfg, cache: Path) -> Path: if hexfont: url = UNSCII_HEX.format(name=cfg.source) else: - url = TOM_THUMB if cfg.source == "tom-thumb" else XORG_RAW.format(name=cfg.source) + url = ( + TOM_THUMB + if cfg.source == "tom-thumb" + else XORG_RAW.format(name=cfg.source) + ) urllib.request.urlretrieve(url, path) return path @@ -84,9 +92,21 @@ def load_font(cfg: FontCfg, cache: Path) -> tuple[dict[int, dict], int]: _HUES = [0.0, 0.08, 0.3, 0.5, 0.62, 0.78] _DARK = [tuple(int(c * 255) for c in colorsys.hls_to_rgb(h, 0.22, 0.95)) for h in _HUES] _PALE = [tuple(int(c * 255) for c in colorsys.hls_to_rgb(h, 0.94, 0.6)) for h in _HUES] -_BRIGHT = [tuple(int(c * 255) for c in colorsys.hls_to_rgb(h, 0.70, 0.95)) for h in _HUES] +_BRIGHT = [ + tuple(int(c * 255) for c in colorsys.hls_to_rgb(h, 0.70, 0.95)) for h in _HUES +] -VARIANTS = ("color", "zebra", "bw", "sent", "dark", "dark-sent", "dim", "sent-dim", "dark-sent-dim") +VARIANTS = ( + "color", + "zebra", + "bw", + "sent", + "dark", + "dark-sent", + "dim", + "sent-dim", + "dark-sent-dim", +) _BLACK = (0, 0, 0) _WHITE = (255, 255, 255) _GRAY = (232, 232, 232) @@ -104,7 +124,9 @@ _STOPWORDS = frozenset( ) -def _row_palette(variant: str, row: int) -> tuple[tuple[int, int, int], tuple[int, int, int]]: +def _row_palette( + variant: str, row: int +) -> tuple[tuple[int, int, int], tuple[int, int, int]]: """(background, default glyph) colors for a row under the given render variant.""" if variant == "color": return _PALE[row % 6], _DARK[row % 6] @@ -155,7 +177,12 @@ def capacity(cfg: FontCfg, size: int = 1568, columns: int = 1) -> tuple[int, int def render( - text: str, cfg: FontCfg, cache: Path, size: int = 1568, variant: str = "color", columns: int = 1 + text: str, + cfg: FontCfg, + cache: Path, + size: int = 1568, + variant: str = "color", + columns: int = 1, ) -> Image.Image: """Fill a size x size grid with `text`; styling per `variant`. @@ -182,8 +209,16 @@ def render( ascent = cfg.ascent if cfg.ascent is not None else font_ascent cols, rows, cap = capacity(cfg, size, columns) text = text[:cap] - sent_idx = _sentence_indices(text) if variant in ("sent", "dark-sent", "sent-dim", "dark-sent-dim") else None - dim_mask = _stopword_mask(text) if variant in ("dim", "sent-dim", "dark-sent-dim") else None + sent_idx = ( + _sentence_indices(text) + if variant in ("sent", "dark-sent", "sent-dim", "dark-sent-dim") + else None + ) + dim_mask = ( + _stopword_mask(text) + if variant in ("dim", "sent-dim", "dark-sent-dim") + else None + ) sent_palette = _BRIGHT if variant in ("dark-sent", "dark-sent-dim") else _DARK dark_bg = variant in ("dark", "dark-sent", "dark-sent-dim") base_color = _BLACK if dark_bg else _WHITE @@ -239,7 +274,9 @@ def render( for y in range(canvas_h): px[x, y] = rule if cfg.native is not None: - img = img.resize((canvas_w * cfg.adv // aw, canvas_h * cfg.pitch // ph), Image.LANCZOS) + img = img.resize( + (canvas_w * cfg.adv // aw, canvas_h * cfg.pitch // ph), Image.LANCZOS + ) if img.size != (size, size): out = Image.new("RGB", (size, size), base_color) out.paste(img, (0, 0)) diff --git a/packages/snapcompact/research/bench_gemini.py b/packages/snapcompact/research/bench_gemini.py index 6c3f636e0..f062ee442 100644 --- a/packages/snapcompact/research/bench_gemini.py +++ b/packages/snapcompact/research/bench_gemini.py @@ -39,7 +39,8 @@ PRICE_IN, PRICE_OUT = 0.6, 4.0 # $/M, matches final.MODELS google/gemini-3.5-fl def _post(body: dict, api_key: str, retries: int = 4) -> dict: payload = json.dumps(body).encode() req = urllib.request.Request( - GEMINI_URL, data=payload, + GEMINI_URL, + data=payload, headers={"content-type": "application/json", "x-goog-api-key": api_key}, ) for attempt in range(retries + 1): @@ -64,7 +65,9 @@ def _post(body: dict, api_key: str, retries: int = 4) -> dict: raise AssertionError("unreachable") -def gemini_complete(api_key: str, blocks: list[dict], resolution: str | None, max_tokens: int) -> dict: +def gemini_complete( + api_key: str, blocks: list[dict], resolution: str | None, max_tokens: int +) -> dict: """blocks: [{"text": str} | {"image_path": Path}]; returns {"text", "usage", "stop"}.""" parts = [] for b in blocks: @@ -74,7 +77,9 @@ def gemini_complete(api_key: str, blocks: list[dict], resolution: str | None, ma part: dict = { "inline_data": { "mime_type": "image/png", - "data": base64.b64encode(Path(b["image_path"]).read_bytes()).decode(), + "data": base64.b64encode( + Path(b["image_path"]).read_bytes() + ).decode(), } } if resolution: @@ -99,7 +104,11 @@ def gemini_complete(api_key: str, blocks: list[dict], resolution: str | None, ma "cache_r": u.get("cachedContentTokenCount", 0), "reasoning": u.get("thoughtsTokenCount", 0), } - stop = "max_tokens" if cand.get("finishReason") == "MAX_TOKENS" else (cand.get("finishReason") or "").lower() + stop = ( + "max_tokens" + if cand.get("finishReason") == "MAX_TOKENS" + else (cand.get("finishReason") or "").lower() + ) return {"text": text, "usage": usage, "stop": stop} @@ -107,8 +116,11 @@ def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--shape-json", required=True) ap.add_argument("--name", required=True) - ap.add_argument("--resolution", default=None, - help="per-part media_resolution level, e.g. MEDIA_RESOLUTION_ULTRA_HIGH; omit for API default") + ap.add_argument( + "--resolution", + default=None, + help="per-part media_resolution level, e.g. MEDIA_RESOLUTION_ULTRA_HIGH; omit for API default", + ) ap.add_argument("--chars", type=int, default=400_000) ap.add_argument("--questions", type=int, default=25) ap.add_argument("--qpb", type=int, default=5) @@ -121,29 +133,49 @@ def main() -> None: api_key = load_env_key("GEMINI_API_KEY", args.env) paras = squad.load_paragraphs(CACHE) flow, offsets = squad.build_flow(paras, args.chars) - questions = squad.sample_chunk_questions(paras, offsets, 0, len(flow), args.questions, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, 0, len(flow), args.questions, args.seed + ) shape, label = json.loads(args.shape_json), args.name cond = f"prod-{label}" size = shape["frameSize"] - frame_dir = CACHE / f"prod-frames-{label}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + frame_dir = ( + CACHE / f"prod-frames-{label}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + ) if not frame_dir.exists() or not any(frame_dir.iterdir()): flow_file = CACHE / f"prod-flow-{sha8(flow)}.txt" flow_file.write_text(flow) subprocess.run( - ["bun", str(HERE / "render_pages.ts"), str(flow_file), json.dumps(shape), str(frame_dir)], + [ + "bun", + str(HERE / "render_pages.ts"), + str(flow_file), + json.dumps(shape), + str(frame_dir), + ], check=True, ) pngs = sorted(frame_dir.glob("page-*.png")) - cols = (size // shape["cellWidth"] - 3) // 2 if shape.get("columns") == 2 else size // shape["cellWidth"] + cols = ( + (size // shape["cellWidth"] - 3) // 2 + if shape.get("columns") == 2 + else size // shape["cellWidth"] + ) rows = size // shape["cellHeight"] // shape.get("lineRepeat", 1) - preamble = load_prompt("qa-image-multi.md").format(k=len(pngs), cols=cols, rows=rows) + preamble = load_prompt("qa-image-multi.md").format( + k=len(pngs), cols=cols, rows=rows + ) if shape.get("columns") == 2: preamble += ( "\nNote: each image lays text out as two word-wrapped newspaper columns separated by a gutter; " "read the left column top to bottom, then the right column." ) - ctx_blocks = [{"text": preamble}, *({"image_path": str(p)} for p in pngs), {"text": "End of images."}] + ctx_blocks = [ + {"text": preamble}, + *({"image_path": str(p)} for p in pngs), + {"text": "End of images."}, + ] out_dir = RESULTS / f"mono-prod-gemini-direct-{label}" out_dir.mkdir(parents=True, exist_ok=True) @@ -153,9 +185,15 @@ def main() -> None: q_block = "\n".join(f"{i + 1}. {q['q']}" for i, q in enumerate(batch)) blocks = [*ctx_blocks, {"text": q_block}] qa = cached( - f"gemini-direct-{GEMINI_MODEL}", "qa-mono-prod-direct", - {"blocks": [{k: str(v) for k, v in blk.items()} for blk in blocks], "resolution": args.resolution}, - lambda blk=blocks: gemini_complete(api_key, blk, args.resolution, args.max_tokens), + f"gemini-direct-{GEMINI_MODEL}", + "qa-mono-prod-direct", + { + "blocks": [{k: str(v) for k, v in blk.items()} for blk in blocks], + "resolution": args.resolution, + }, + lambda blk=blocks: gemini_complete( + api_key, blk, args.resolution, args.max_tokens + ), args.fresh, ) answers.extend(squad.parse_numbered(qa["text"], len(batch))) @@ -163,21 +201,39 @@ def main() -> None: stops.append(qa["stop"]) rows_out = [ { - "model": GEMINI_MODEL, "cond": cond, "pos_rel": q["pos_rel"], "q": q["q"], "answer": a, - "golds": q["golds"], "em": squad.exact_match(a, q["golds"]), "f1": squad.f1(a, q["golds"]), + "model": GEMINI_MODEL, + "cond": cond, + "pos_rel": q["pos_rel"], + "q": q["q"], + "answer": a, + "golds": q["golds"], + "em": squad.exact_match(a, q["golds"]), + "f1": squad.f1(a, q["golds"]), "abstained": "unreadable" in a.lower(), } for q, a in zip(questions, answers) ] - u = {k: sum(x[k] for x in usages) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} + u = { + k: sum(x[k] for x in usages) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } cost = (u["in"] + 0.1 * u["cache_r"]) / 1e6 * PRICE_IN + u["out"] / 1e6 * PRICE_OUT summary = { - "cond": cond, "n": len(rows_out), "imgs": len(pngs), "resolution": args.resolution, + "cond": cond, + "n": len(rows_out), + "imgs": len(pngs), + "resolution": args.resolution, "em": sum(r["em"] for r in rows_out) / len(rows_out), "f1": sum(r["f1"] for r in rows_out) / len(rows_out), "abst": sum(r["abstained"] for r in rows_out), - "tok_in": u["in"], "tok_cached": u["cache_r"], "tok_out": u["out"], "reas": u["reasoning"], - "cost": cost, "stop": next((s for s in stops if s == "max_tokens"), stops[-1] if stops else ""), + "tok_in": u["in"], + "tok_cached": u["cache_r"], + "tok_out": u["out"], + "reas": u["reasoning"], + "cost": cost, + "stop": next( + (s for s in stops if s == "max_tokens"), stops[-1] if stops else "" + ), } (out_dir / "records.jsonl").write_text("\n".join(json.dumps(r) for r in rows_out)) (out_dir / "summary.json").write_text(json.dumps([summary], indent=1)) diff --git a/packages/snapcompact/research/bench_kimi.py b/packages/snapcompact/research/bench_kimi.py index be14460ac..5d020a6db 100644 --- a/packages/snapcompact/research/bench_kimi.py +++ b/packages/snapcompact/research/bench_kimi.py @@ -42,16 +42,25 @@ def fireworks_complete(messages: list[dict], max_tokens: int) -> tuple[str, dict return [ {"type": "text", "text": b["text"]} if "text" in b - else {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{_png_b64(b['image_path'])}"}} + else { + "type": "image_url", + "image_url": { + "url": f"data:image/png;base64,{_png_b64(b['image_path'])}" + }, + } for b in blocks ] body = { "model": FW_MODEL, - "messages": [{"role": m["role"], "content": content(m["content"])} for m in messages], + "messages": [ + {"role": m["role"], "content": content(m["content"])} for m in messages + ], "max_tokens": max_tokens, } - out = _post(FW_URL, body, {"authorization": f"Bearer {load_env_key('FIREWORKS_API_KEY')}"}) + out = _post( + FW_URL, body, {"authorization": f"Bearer {load_env_key('FIREWORKS_API_KEY')}"} + ) choice = (out.get("choices") or [{}])[0] text = (choice.get("message") or {}).get("content") or "" if isinstance(text, list): @@ -62,9 +71,15 @@ def fireworks_complete(messages: list[dict], max_tokens: int) -> tuple[str, dict "out": u.get("completion_tokens", 0), "cache_w": 0, "cache_r": 0, - "reasoning": (u.get("completion_tokens_details") or {}).get("reasoning_tokens", 0), + "reasoning": (u.get("completion_tokens_details") or {}).get( + "reasoning_tokens", 0 + ), } - stop = "max_tokens" if choice.get("finish_reason") == "length" else (choice.get("finish_reason") or "") + stop = ( + "max_tokens" + if choice.get("finish_reason") == "length" + else (choice.get("finish_reason") or "") + ) return text, usage, stop @@ -81,28 +96,50 @@ def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--shape-json", required=True) ap.add_argument("--name", required=True) - ap.add_argument("--route", default="openrouter", choices=["openrouter", "fireworks"]) + ap.add_argument( + "--route", default="openrouter", choices=["openrouter", "fireworks"] + ) ap.add_argument("--chars", type=int, default=400_000) ap.add_argument("--questions", type=int, default=25) ap.add_argument("--seed", type=int, default=42) ap.add_argument("--max-tokens", type=int, default=32768) - ap.add_argument("--frame-tokens", type=float, default=None, help="expected billed tokens per frame (bill sanity)") + ap.add_argument( + "--frame-tokens", + type=float, + default=None, + help="expected billed tokens per frame (bill sanity)", + ) ap.add_argument("--fresh", action="store_true") args = ap.parse_args() shape, label = json.loads(args.shape_json), args.name - keys = {"openrouter": load_env_key("OPENROUTER_API_KEY"), "anthropic": "", "openai": ""} + keys = { + "openrouter": load_env_key("OPENROUTER_API_KEY"), + "anthropic": "", + "openai": "", + } paras = squad.load_paragraphs(CACHE) flow, offsets = squad.build_flow(paras, args.chars) - questions = squad.sample_chunk_questions(paras, offsets, 0, len(flow), args.questions, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, 0, len(flow), args.questions, args.seed + ) # Production frames, same dir-keying convention as mono_prod.py (reuses its renders). - frame_dir = CACHE / f"prod-frames-{label}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + frame_dir = ( + CACHE / f"prod-frames-{label}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + ) if not frame_dir.exists() or not any(frame_dir.iterdir()): flow_file = CACHE / f"prod-flow-{sha8(flow)}.txt" flow_file.write_text(flow) subprocess.run( - ["bun", str(HERE / "render_pages.ts"), str(flow_file), json.dumps(shape), str(frame_dir)], check=True + [ + "bun", + str(HERE / "render_pages.ts"), + str(flow_file), + json.dumps(shape), + str(frame_dir), + ], + check=True, ) pngs = sorted(frame_dir.glob("page-*.png")) n_frames = len(pngs) @@ -130,37 +167,63 @@ def main() -> None: continue chunk_pngs = pngs[lo:hi] frames_sent += len(chunk_pngs) - preamble = load_prompt("qa-image-multi.md").format(k=len(chunk_pngs), cols=cols, rows=rows) - ctx = [{"text": preamble}, *({"image_path": p} for p in chunk_pngs), {"text": "End of images.", "cache": True}] + preamble = load_prompt("qa-image-multi.md").format( + k=len(chunk_pngs), cols=cols, rows=rows + ) + ctx = [ + {"text": preamble}, + *({"image_path": p} for p in chunk_pngs), + {"text": "End of images.", "cache": True}, + ] q_block = "\n".join(f"{i + 1}. {q['q']}" for i, q in enumerate(qs)) messages = [{"role": "user", "content": [*ctx, {"text": q_block}]}] - tag = "qa-mono-prod-chunk8" if args.route == "openrouter" else "qa-mono-prod-chunk8-fw" + tag = ( + "qa-mono-prod-chunk8" + if args.route == "openrouter" + else "qa-mono-prod-chunk8-fw" + ) if args.route == "openrouter": fn = lambda m=messages: dict( # noqa: E731 - zip(("text", "usage", "stop"), llm_complete(keys, OR_MODEL, m, max_tokens=args.max_tokens, effort=None)) + zip( + ("text", "usage", "stop"), + llm_complete( + keys, OR_MODEL, m, max_tokens=args.max_tokens, effort=None + ), + ) ) else: fn = lambda m=messages: dict( # noqa: E731 zip(("text", "usage", "stop"), fireworks_complete(m, args.max_tokens)) ) - qa = cached(OR_MODEL, tag, {"messages": messages, "effort": None}, fn, args.fresh) + qa = cached( + OR_MODEL, tag, {"messages": messages, "effort": None}, fn, args.fresh + ) for q, a in zip(qs, squad.parse_numbered(qa["text"], len(qs))): answers_by_q[q["q"]] = a usages.append(qa["usage"]) stops.append(qa["stop"]) - print(f" chunk {ci} frames[{lo}:{hi}] nq={len(qs)} in={qa['usage']['in']} stop={qa['stop']}") + print( + f" chunk {ci} frames[{lo}:{hi}] nq={len(qs)} in={qa['usage']['in']} stop={qa['stop']}" + ) rows_out = [ { - "model": OR_MODEL, "cond": f"bench-{label}", "pos_rel": q["pos_rel"], "q": q["q"], - "answer": answers_by_q.get(q["q"], ""), "golds": q["golds"], + "model": OR_MODEL, + "cond": f"bench-{label}", + "pos_rel": q["pos_rel"], + "q": q["q"], + "answer": answers_by_q.get(q["q"], ""), + "golds": q["golds"], "em": squad.exact_match(answers_by_q.get(q["q"], ""), q["golds"]), "f1": squad.f1(answers_by_q.get(q["q"], ""), q["golds"]), "abstained": "unreadable" in answers_by_q.get(q["q"], "").lower(), } for q in questions ] - u = {k: sum(x[k] for x in usages) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} + u = { + k: sum(x[k] for x in usages) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } price_in, price_out = MODELS[OR_MODEL] if args.route == "openrouter" else FW_PRICE cost = u["in"] / 1e6 * price_in + u["out"] / 1e6 * price_out quart = [] @@ -170,16 +233,30 @@ def main() -> None: expect_frame = args.frame_tokens or (math.ceil(size / 28) ** 2 + 5) bill_ratio = u["in"] / (frames_sent * expect_frame) if frames_sent else float("nan") summary = { - "cond": f"bench-{label}", "route": args.route, "n": len(rows_out), "imgs": n_frames, - "frames_sent": frames_sent, "chunks": chunks, + "cond": f"bench-{label}", + "route": args.route, + "n": len(rows_out), + "imgs": n_frames, + "frames_sent": frames_sent, + "chunks": chunks, "em": sum(r["em"] for r in rows_out) / len(rows_out), "f1": sum(r["f1"] for r in rows_out) / len(rows_out), "abst": sum(r["abstained"] for r in rows_out), - "tok_in": u["in"], "tok_out": u["out"], "reas": u["reasoning"], "cost": cost, - "expect_frame_tokens": expect_frame, "bill_ratio": round(bill_ratio, 3), - "chars_per_dollar": args.chars / cost, "chars_per_mtok_in": args.chars / u["in"] * 1e6, - "stop": next((s for s in stops if s == "max_tokens"), stops[-1] if stops else ""), - "q1": quart[0], "q2": quart[1], "q3": quart[2], "q4": quart[3], + "tok_in": u["in"], + "tok_out": u["out"], + "reas": u["reasoning"], + "cost": cost, + "expect_frame_tokens": expect_frame, + "bill_ratio": round(bill_ratio, 3), + "chars_per_dollar": args.chars / cost, + "chars_per_mtok_in": args.chars / u["in"] * 1e6, + "stop": next( + (s for s in stops if s == "max_tokens"), stops[-1] if stops else "" + ), + "q1": quart[0], + "q2": quart[1], + "q3": quart[2], + "q4": quart[3], } out_dir = RESULTS / f"bench-kimi-{label}" out_dir.mkdir(parents=True, exist_ok=True) @@ -190,7 +267,9 @@ def main() -> None: f"tok_in={u['in']} bill_ratio={summary['bill_ratio']} ${cost:.3f} " f"chars/$={summary['chars_per_dollar']:.0f}" ) - print("F1 by quartile: " + " ".join(f"q{i + 1}={v:.3f}" for i, v in enumerate(quart))) + print( + "F1 by quartile: " + " ".join(f"q{i + 1}={v:.3f}" for i, v in enumerate(quart)) + ) if __name__ == "__main__": diff --git a/packages/snapcompact/research/bench_kimi_probe.py b/packages/snapcompact/research/bench_kimi_probe.py index 73409d26b..e383edf6a 100644 --- a/packages/snapcompact/research/bench_kimi_probe.py +++ b/packages/snapcompact/research/bench_kimi_probe.py @@ -25,7 +25,11 @@ from providers import _png_b64, _post, load_env_key # noqa: E402 from run import CACHE # noqa: E402 ROUTES = { - "openrouter": ("https://openrouter.ai/api/v1/chat/completions", "OPENROUTER_API_KEY", "moonshotai/kimi-k2.6"), + "openrouter": ( + "https://openrouter.ai/api/v1/chat/completions", + "OPENROUTER_API_KEY", + "moonshotai/kimi-k2.6", + ), "fireworks": ( "https://api.fireworks.ai/inference/v1/chat/completions", "FIREWORKS_API_KEY", @@ -55,10 +59,22 @@ def frame_at(px: int) -> Path: def ask(route: str, content: list[dict]) -> dict: url, key_var, model = ROUTES[route] - body = {"model": model, "messages": [{"role": "user", "content": content}], "max_tokens": 64} - out = _post(url, body, {"authorization": f"Bearer {load_env_key(key_var)}", "user-agent": "bench/1.0"}) + body = { + "model": model, + "messages": [{"role": "user", "content": content}], + "max_tokens": 64, + } + out = _post( + url, + body, + {"authorization": f"Bearer {load_env_key(key_var)}", "user-agent": "bench/1.0"}, + ) choice = (out.get("choices") or [{}])[0] - return {"usage": out.get("usage", {}), "provider": out.get("provider"), "finish": choice.get("finish_reason")} + return { + "usage": out.get("usage", {}), + "provider": out.get("provider"), + "finish": choice.get("finish_reason"), + } def main() -> None: @@ -71,7 +87,10 @@ def main() -> None: base = ask(args.route, [tail])["usage"].get("prompt_tokens", 0) print(f"route={args.route} text-only prompt_tokens={base}") for px in (int(s) for s in args.sizes.split(",")): - img = {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{_png_b64(frame_at(px))}"}} + img = { + "type": "image_url", + "image_url": {"url": f"data:image/png;base64,{_png_b64(frame_at(px))}"}, + } try: r = ask(args.route, [img, tail]) except SystemExit as e: diff --git a/packages/snapcompact/research/diag_glm_forensics.py b/packages/snapcompact/research/diag_glm_forensics.py index 6274bb3c0..6098532d3 100644 --- a/packages/snapcompact/research/diag_glm_forensics.py +++ b/packages/snapcompact/research/diag_glm_forensics.py @@ -26,14 +26,25 @@ from run import CACHE, QA_CACHE, load_prompt, sha8 # noqa: E402 def build_batches(shape_name: str, chars: int, n_questions: int, qpb: int, seed: int): paras = squad.load_paragraphs(CACHE) flow, offsets = squad.build_flow(paras, chars) - questions = squad.sample_chunk_questions(paras, offsets, 0, len(flow), n_questions, seed) + questions = squad.sample_chunk_questions( + paras, offsets, 0, len(flow), n_questions, seed + ) shape = SHAPES[shape_name] - frame_dir = CACHE / f"prod-frames-{shape_name}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + frame_dir = ( + CACHE + / f"prod-frames-{shape_name}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + ) pngs = sorted(frame_dir.glob("page-*.png")) repeat = shape.get("lineRepeat", 1) - cols = (SIZE // shape["cellWidth"] - 3) // 2 if shape.get("columns") == 2 else SIZE // shape["cellWidth"] + cols = ( + (SIZE // shape["cellWidth"] - 3) // 2 + if shape.get("columns") == 2 + else SIZE // shape["cellWidth"] + ) rows = SIZE // shape["cellHeight"] // repeat - preamble = load_prompt("qa-image-multi.md").format(k=len(pngs), cols=cols, rows=rows) + preamble = load_prompt("qa-image-multi.md").format( + k=len(pngs), cols=cols, rows=rows + ) if shape.get("columns") == 2: preamble += ( "\nNote: each image lays text out as two word-wrapped newspaper columns separated by a gutter; " @@ -45,7 +56,11 @@ def build_batches(shape_name: str, chars: int, n_questions: int, qpb: int, seed: "background, then repeated on a pale highlight band. The copies show identical characters; " "cross-check between them when a glyph is hard to read, and do not treat copies as separate text." ) - ctx_blocks = [{"text": preamble}, *({"image_path": p} for p in pngs), {"text": "End of images.", "cache": True}] + ctx_blocks = [ + {"text": preamble}, + *({"image_path": p} for p in pngs), + {"text": "End of images.", "cache": True}, + ] batches = [] for b in range(0, len(questions), qpb): batch = questions[b : b + qpb] @@ -66,11 +81,15 @@ def main() -> None: ap.add_argument("--effort", default=None) args = ap.parse_args() - pngs, batches = build_batches(args.shape, args.chars, args.questions, args.qpb, args.seed) + pngs, batches = build_batches( + args.shape, args.chars, args.questions, args.qpb, args.seed + ) print(f"frames={len(pngs)} batches={len(batches)}") for i, (batch, messages) in enumerate(batches): payload = {"messages": messages, "effort": args.effort} - key = sha8(args.model, "qa-mono-prod", json.dumps(payload, sort_keys=True, default=str)) + key = sha8( + args.model, "qa-mono-prod", json.dumps(payload, sort_keys=True, default=str) + ) path = QA_CACHE / f"{key}.json" print(f"\n=== batch {i + 1} key={key} cached={path.exists()} ===") for j, q in enumerate(batch): diff --git a/packages/snapcompact/research/diag_glm_mono.py b/packages/snapcompact/research/diag_glm_mono.py index 93b336c28..11c6de589 100644 --- a/packages/snapcompact/research/diag_glm_mono.py +++ b/packages/snapcompact/research/diag_glm_mono.py @@ -53,24 +53,37 @@ def to_chat(messages: list[dict]) -> list[dict]: return out -def score(questions, answers, label: str, out_dir: Path | None = None, extra: dict | None = None): +def score( + questions, + answers, + label: str, + out_dir: Path | None = None, + extra: dict | None = None, +): rows = [ { - "q": q["q"], "pos_rel": q["pos_rel"], "answer": a, "golds": q["golds"], - "em": squad.exact_match(a, q["golds"]), "f1": squad.f1(a, q["golds"]), + "q": q["q"], + "pos_rel": q["pos_rel"], + "answer": a, + "golds": q["golds"], + "em": squad.exact_match(a, q["golds"]), + "f1": squad.f1(a, q["golds"]), "abstained": "unreadable" in a.lower(), } for q, a in zip(questions, answers) ] n = len(rows) summary = { - "label": label, "n": n, + "label": label, + "n": n, "em": sum(r["em"] for r in rows) / n, "f1": sum(r["f1"] for r in rows) / n, "abst": sum(r["abstained"] for r in rows), **(extra or {}), } - print(f"{label:<34} n={n} f1={summary['f1']:.3f} em={summary['em']:.3f} abst={summary['abst']}") + print( + f"{label:<34} n={n} f1={summary['f1']:.3f} em={summary['em']:.3f} abst={summary['abst']}" + ) if out_dir: out_dir.mkdir(parents=True, exist_ok=True) (out_dir / "records.jsonl").write_text("\n".join(json.dumps(r) for r in rows)) @@ -80,7 +93,9 @@ def score(questions, answers, label: str, out_dir: Path | None = None, extra: di def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--route", required=True, choices=["openrouter", "zai", "rescore-cache"]) + ap.add_argument( + "--route", required=True, choices=["openrouter", "zai", "rescore-cache"] + ) ap.add_argument("--shape", default="8on16-bw") ap.add_argument("--chars", type=int, default=400_000) ap.add_argument("--questions", type=int, default=25) @@ -92,7 +107,9 @@ def main() -> None: ap.add_argument("--env", default="~/.env") args = ap.parse_args() - pngs, batches = build_batches(args.shape, args.chars, args.questions, args.qpb, args.seed) + pngs, batches = build_batches( + args.shape, args.chars, args.questions, args.qpb, args.seed + ) questions = [q for batch, _ in batches for q in batch] if args.route == "rescore-cache": @@ -100,12 +117,20 @@ def main() -> None: answers = [] for batch, messages in batches: payload = {"messages": messages, "effort": None} - key = sha8(args.model_or, "qa-mono-prod", json.dumps(payload, sort_keys=True, default=str)) + key = sha8( + args.model_or, + "qa-mono-prod", + json.dumps(payload, sort_keys=True, default=str), + ) path = QA_CACHE / f"{key}.json" text = json.loads(path.read_text())["text"] if path.exists() else "" answers.extend(parse_robust(text, len(batch))) - score(questions, answers, f"openrouter-cached-robust-{args.shape}", - RESULTS / f"diag-glm-or-robust-{args.shape}") + score( + questions, + answers, + f"openrouter-cached-robust-{args.shape}", + RESULTS / f"diag-glm-or-robust-{args.shape}", + ) return keys = { @@ -116,7 +141,9 @@ def main() -> None: def run_batch(item): batch, messages = item chat = to_chat(messages) - r = complete(args.route, keys, chat, args.model_or, args.model_zai, args.max_tokens) + r = complete( + args.route, keys, chat, args.model_or, args.model_zai, args.max_tokens + ) if "http_error" in r: print(f" HTTP {r['http_error']}: {r['body'][:200]}") return [""] * len(batch), {} @@ -128,7 +155,9 @@ def main() -> None: tok_in = sum(u.get("prompt_tokens", 0) for _, u in results) tok_out = sum(u.get("completion_tokens", 0) for _, u in results) score( - questions, answers, f"{args.route}-{args.shape}-{len(pngs)}f", + questions, + answers, + f"{args.route}-{args.shape}-{len(pngs)}f", RESULTS / f"diag-glm-{args.route}-{args.shape}-{len(pngs)}f", {"imgs": len(pngs), "tok_in": tok_in, "tok_out": tok_out}, ) diff --git a/packages/snapcompact/research/diag_glm_probe.py b/packages/snapcompact/research/diag_glm_probe.py index d48aa7e6c..eada59a12 100644 --- a/packages/snapcompact/research/diag_glm_probe.py +++ b/packages/snapcompact/research/diag_glm_probe.py @@ -52,7 +52,14 @@ def img_block(path: Path) -> dict: return {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{b64}"}} -def complete(route: str, keys: dict, messages: list[dict], model_or: str, model_zai: str, max_tokens: int) -> dict: +def complete( + route: str, + keys: dict, + messages: list[dict], + model_or: str, + model_zai: str, + max_tokens: int, +) -> dict: if route == "zai": body = {"model": model_zai, "messages": messages, "max_tokens": max_tokens} out = post(ZAI_URL, body, keys["zai"]) @@ -68,7 +75,9 @@ def complete(route: str, keys: dict, messages: list[dict], model_or: str, model_ text = "".join(p.get("text", "") for p in text if isinstance(p, dict)) return { "text": text, - "reasoning_text_len": len(msg.get("reasoning_content") or msg.get("reasoning") or ""), + "reasoning_text_len": len( + msg.get("reasoning_content") or msg.get("reasoning") or "" + ), "usage": out.get("usage", {}), "finish": choice.get("finish_reason"), "provider": out.get("provider"), @@ -80,7 +89,10 @@ def frames_for(shape_name: str, chars: int) -> tuple[list[Path], str, int]: paras = squad.load_paragraphs(CACHE) flow, _ = squad.build_flow(paras, chars) shape = SHAPES[shape_name] - frame_dir = CACHE / f"prod-frames-{shape_name}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + frame_dir = ( + CACHE + / f"prod-frames-{shape_name}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + ) pngs = sorted(frame_dir.glob("page-*.png")) assert pngs, f"no frames in {frame_dir}; run mono_prod first" return pngs, flow, len(flow) @@ -109,7 +121,12 @@ def main() -> None: if args.mode == "smoke": for route in routes: - msgs = [{"role": "user", "content": [{"type": "text", "text": "Reply with exactly: PONG"}]}] + msgs = [ + { + "role": "user", + "content": [{"type": "text", "text": "Reply with exactly: PONG"}], + } + ] r = complete(route, keys, msgs, args.model_or, args.model_zai, 1024) print(f"[{route}] {json.dumps(r, default=str)[:600]}") return @@ -119,15 +136,22 @@ def main() -> None: if args.mode == "bill": for n in (int(x) for x in args.counts.split(",")): - blocks = [{"type": "text", "text": f"You will see {n} images. Reply with exactly: OK"}] + blocks = [ + { + "type": "text", + "text": f"You will see {n} images. Reply with exactly: OK", + } + ] blocks += [img_block(p) for p in pngs[:n]] msgs = [{"role": "user", "content": blocks}] for route in routes: r = complete(route, keys, msgs, args.model_or, args.model_zai, 2048) u = r.get("usage", {}) pt = u.get("prompt_tokens", 0) - print(f"[{route}] n={n:>2} prompt_tokens={pt:>7} per_img={(pt / max(n, 1)):>8.1f} " - f"finish={r.get('finish')} provider={r.get('provider')} text={r.get('text', '')[:40]!r}") + print( + f"[{route}] n={n:>2} prompt_tokens={pt:>7} per_img={(pt / max(n, 1)):>8.1f} " + f"finish={r.get('finish')} provider={r.get('provider')} text={r.get('text', '')[:40]!r}" + ) return if args.mode == "frame": @@ -135,11 +159,17 @@ def main() -> None: ask = "Transcribe the first 5 text lines of this image exactly." if args.question: ask += f"\nThen answer from the image text: {args.question}\nFormat: TRANSCRIPT lines, then ANSWER: ." - msgs = [{"role": "user", "content": [{"type": "text", "text": ask}, img_block(p)]}] + msgs = [ + {"role": "user", "content": [{"type": "text", "text": ask}, img_block(p)]} + ] for route in routes: - r = complete(route, keys, msgs, args.model_or, args.model_zai, args.max_tokens) + r = complete( + route, keys, msgs, args.model_or, args.model_zai, args.max_tokens + ) u = r.get("usage", {}) - print(f"\n[{route}] frame={p.name} prompt_tokens={u.get('prompt_tokens')} finish={r.get('finish')}") + print( + f"\n[{route}] frame={p.name} prompt_tokens={u.get('prompt_tokens')} finish={r.get('finish')}" + ) print("\n".join(" | " + ln for ln in r.get("text", "").splitlines()[:12])) return @@ -150,23 +180,34 @@ def main() -> None: questions = squad.sample_chunk_questions(paras, offsets, 0, len(flow2), 25, 42) batch = questions[15:20] from mono_prod import SIZE + shape = SHAPES[args.shape] cols = SIZE // shape["cellWidth"] rows = SIZE // shape["cellHeight"] from run import load_prompt - preamble = load_prompt("qa-image-multi.md").format(k=len(pngs), cols=cols, rows=rows) + + preamble = load_prompt("qa-image-multi.md").format( + k=len(pngs), cols=cols, rows=rows + ) q_block = "\n".join(f"{i + 1}. {q['q']}" for i, q in enumerate(batch)) blocks = [{"type": "text", "text": preamble}] blocks += [img_block(p) for p in pngs] - blocks += [{"type": "text", "text": "End of images."}, {"type": "text", "text": q_block}] + blocks += [ + {"type": "text", "text": "End of images."}, + {"type": "text", "text": q_block}, + ] msgs = [{"role": "user", "content": blocks}] for i, q in enumerate(batch): print(f"Q{i + 1}: {q['q']} golds={q['golds']}") for route in routes: - r = complete(route, keys, msgs, args.model_or, args.model_zai, args.max_tokens) + r = complete( + route, keys, msgs, args.model_or, args.model_zai, args.max_tokens + ) u = r.get("usage", {}) - print(f"\n[{route}] prompt_tokens={u.get('prompt_tokens')} completion={u.get('completion_tokens')} " - f"finish={r.get('finish')} provider={r.get('provider')}") + print( + f"\n[{route}] prompt_tokens={u.get('prompt_tokens')} completion={u.get('completion_tokens')} " + f"finish={r.get('finish')} provider={r.get('provider')}" + ) print("\n".join(" | " + ln for ln in r.get("text", "").splitlines()[:10])) diff --git a/packages/snapcompact/research/diag_kimi_chunked.py b/packages/snapcompact/research/diag_kimi_chunked.py index 47bf91d87..f44c486e3 100644 --- a/packages/snapcompact/research/diag_kimi_chunked.py +++ b/packages/snapcompact/research/diag_kimi_chunked.py @@ -31,12 +31,19 @@ CHUNKS = [(0, 8), (6, 14), (13, 21)] def main() -> None: - keys = {"openrouter": load_env_key("OPENROUTER_API_KEY"), "anthropic": "", "openai": ""} + keys = { + "openrouter": load_env_key("OPENROUTER_API_KEY"), + "anthropic": "", + "openai": "", + } paras = squad.load_paragraphs(CACHE) flow, offsets = squad.build_flow(paras, 400_000) questions = squad.sample_chunk_questions(paras, offsets, 0, len(flow), 25, 42) shape = SHAPES[SHAPE_NAME] - frame_dir = CACHE / f"prod-frames-{SHAPE_NAME}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + frame_dir = ( + CACHE + / f"prod-frames-{SHAPE_NAME}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + ) pngs = sorted(frame_dir.glob("page-*.png")) assert len(pngs) == 21, len(pngs) cols = SIZE // shape["cellWidth"] @@ -63,14 +70,25 @@ def main() -> None: if not qs: continue chunk_pngs = pngs[lo:hi] - preamble = load_prompt("qa-image-multi.md").format(k=len(chunk_pngs), cols=cols, rows=rows) - ctx = [{"text": preamble}, *({"image_path": p} for p in chunk_pngs), {"text": "End of images.", "cache": True}] + preamble = load_prompt("qa-image-multi.md").format( + k=len(chunk_pngs), cols=cols, rows=rows + ) + ctx = [ + {"text": preamble}, + *({"image_path": p} for p in chunk_pngs), + {"text": "End of images.", "cache": True}, + ] q_block = "\n".join(f"{i + 1}. {q['q']}" for i, q in enumerate(qs)) messages = [{"role": "user", "content": [*ctx, {"text": q_block}]}] qa = cached( - MODEL, "qa-mono-prod-chunk8", {"messages": messages, "effort": None}, + MODEL, + "qa-mono-prod-chunk8", + {"messages": messages, "effort": None}, lambda m=messages: dict( - zip(("text", "usage", "stop"), llm_complete(keys, MODEL, m, max_tokens=32768, effort=None)) + zip( + ("text", "usage", "stop"), + llm_complete(keys, MODEL, m, max_tokens=32768, effort=None), + ) ), False, ) @@ -78,19 +96,28 @@ def main() -> None: answers_by_q[q["q"]] = a usages.append(qa["usage"]) stops.append(qa["stop"]) - print(f"chunk {ci} frames[{lo}:{hi}] nq={len(qs)} in={qa['usage']['in']} stop={qa['stop']}") + print( + f"chunk {ci} frames[{lo}:{hi}] nq={len(qs)} in={qa['usage']['in']} stop={qa['stop']}" + ) rows_out = [ { - "model": MODEL, "cond": "diag-chunk8-8on16-bw", "pos_rel": q["pos_rel"], "q": q["q"], - "answer": answers_by_q.get(q["q"], ""), "golds": q["golds"], + "model": MODEL, + "cond": "diag-chunk8-8on16-bw", + "pos_rel": q["pos_rel"], + "q": q["q"], + "answer": answers_by_q.get(q["q"], ""), + "golds": q["golds"], "em": squad.exact_match(answers_by_q.get(q["q"], ""), q["golds"]), "f1": squad.f1(answers_by_q.get(q["q"], ""), q["golds"]), "abstained": "unreadable" in answers_by_q.get(q["q"], "").lower(), } for q in questions ] - u = {k: sum(x[k] for x in usages) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} + u = { + k: sum(x[k] for x in usages) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } price_in, price_out = MODELS[MODEL] cost = u["in"] / 1e6 * price_in + u["out"] / 1e6 * price_out quart = [] @@ -98,20 +125,34 @@ def main() -> None: sel = [r["f1"] for r in rows_out if lo <= r["pos_rel"] < hi] quart.append(sum(sel) / len(sel) if sel else float("nan")) summary = { - "cond": "diag-chunk8-8on16-bw", "n": len(rows_out), "imgs": 21, + "cond": "diag-chunk8-8on16-bw", + "n": len(rows_out), + "imgs": 21, "em": sum(r["em"] for r in rows_out) / len(rows_out), "f1": sum(r["f1"] for r in rows_out) / len(rows_out), "abst": sum(r["abstained"] for r in rows_out), - "tok_in": u["in"], "tok_out": u["out"], "reas": u["reasoning"], "cost": cost, - "stop": next((s for s in stops if s == "max_tokens"), stops[-1] if stops else ""), - "q1": quart[0], "q2": quart[1], "q3": quart[2], "q4": quart[3], + "tok_in": u["in"], + "tok_out": u["out"], + "reas": u["reasoning"], + "cost": cost, + "stop": next( + (s for s in stops if s == "max_tokens"), stops[-1] if stops else "" + ), + "q1": quart[0], + "q2": quart[1], + "q3": quart[2], + "q4": quart[3], } out_dir = RESULTS / "diag-kimi-chunk8-8on16-bw" out_dir.mkdir(parents=True, exist_ok=True) (out_dir / "records.jsonl").write_text("\n".join(json.dumps(r) for r in rows_out)) (out_dir / "summary.json").write_text(json.dumps([summary], indent=1)) - print(f"diag-chunk8 f1={summary['f1']:.3f} em={summary['em']:.3f} abst={summary['abst']} ${summary['cost']:.2f}") - print("F1 by quartile: " + " ".join(f"q{i + 1}={v:.3f}" for i, v in enumerate(quart))) + print( + f"diag-chunk8 f1={summary['f1']:.3f} em={summary['em']:.3f} abst={summary['abst']} ${summary['cost']:.2f}" + ) + print( + "F1 by quartile: " + " ".join(f"q{i + 1}={v:.3f}" for i, v in enumerate(quart)) + ) if __name__ == "__main__": diff --git a/packages/snapcompact/research/diag_kimi_forensics.py b/packages/snapcompact/research/diag_kimi_forensics.py index 102618a4f..6e7057aad 100644 --- a/packages/snapcompact/research/diag_kimi_forensics.py +++ b/packages/snapcompact/research/diag_kimi_forensics.py @@ -23,17 +23,32 @@ QA_CACHE = CACHE / "qa" MODEL = "moonshotai/kimi-k2.6" -def build_messages(shape_name: str, chars: int = 400_000, questions: int = 25, qpb: int = 5, seed: int = 42): +def build_messages( + shape_name: str, + chars: int = 400_000, + questions: int = 25, + qpb: int = 5, + seed: int = 42, +): paras = squad.load_paragraphs(CACHE) flow, offsets = squad.build_flow(paras, chars) qs = squad.sample_chunk_questions(paras, offsets, 0, len(flow), questions, seed) shape = SHAPES[shape_name] - frame_dir = CACHE / f"prod-frames-{shape_name}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + frame_dir = ( + CACHE + / f"prod-frames-{shape_name}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + ) pngs = sorted(frame_dir.glob("page-*.png")) repeat = shape.get("lineRepeat", 1) - cols = (SIZE // shape["cellWidth"] - 3) // 2 if shape.get("columns") == 2 else SIZE // shape["cellWidth"] + cols = ( + (SIZE // shape["cellWidth"] - 3) // 2 + if shape.get("columns") == 2 + else SIZE // shape["cellWidth"] + ) rows = SIZE // shape["cellHeight"] // repeat - preamble = load_prompt("qa-image-multi.md").format(k=len(pngs), cols=cols, rows=rows) + preamble = load_prompt("qa-image-multi.md").format( + k=len(pngs), cols=cols, rows=rows + ) if shape.get("columns") == 2: preamble += ( "\nNote: each image lays text out as two word-wrapped newspaper columns separated by a gutter; " @@ -45,7 +60,11 @@ def build_messages(shape_name: str, chars: int = 400_000, questions: int = 25, q "background, then repeated on a pale highlight band. The copies show identical characters; " "cross-check between them when a glyph is hard to read, and do not treat copies as separate text." ) - ctx_blocks = [{"text": preamble}, *({"image_path": p} for p in pngs), {"text": "End of images.", "cache": True}] + ctx_blocks = [ + {"text": preamble}, + *({"image_path": p} for p in pngs), + {"text": "End of images.", "cache": True}, + ] batches = [] for b in range(0, len(qs), qpb): batch = qs[b : b + qpb] @@ -60,7 +79,13 @@ def main() -> None: pngs, batches = build_messages(shape_name) print(f"\n=== {shape_name}: {len(pngs)} frames ===") for bi, (batch, messages) in enumerate(batches): - key = sha8(MODEL, "qa-mono-prod", json.dumps({"messages": messages, "effort": None}, sort_keys=True, default=str)) + key = sha8( + MODEL, + "qa-mono-prod", + json.dumps( + {"messages": messages, "effort": None}, sort_keys=True, default=str + ), + ) path = QA_CACHE / f"{key}.json" if not path.exists(): print(f"--- batch {bi}: cache MISS ({key})") diff --git a/packages/snapcompact/research/diag_kimi_mono.py b/packages/snapcompact/research/diag_kimi_mono.py index 76fa2ae7c..3e9216581 100644 --- a/packages/snapcompact/research/diag_kimi_mono.py +++ b/packages/snapcompact/research/diag_kimi_mono.py @@ -38,12 +38,23 @@ def fw_complete(messages: list[dict], max_tokens: int = 32768) -> dict: if "text" in b: out.append({"type": "text", "text": b["text"]}) else: - out.append({"type": "image_url", "image_url": {"url": f"data:image/png;base64,{_png_b64(b['image_path'])}"}}) + out.append( + { + "type": "image_url", + "image_url": { + "url": f"data:image/png;base64,{_png_b64(b['image_path'])}" + }, + } + ) return out - chat_messages = [{"role": m["role"], "content": content(m["content"])} for m in messages] + chat_messages = [ + {"role": m["role"], "content": content(m["content"])} for m in messages + ] body = {"model": FW_MODEL, "messages": chat_messages, "max_tokens": max_tokens} - out = _post(FW_URL, body, {"authorization": f"Bearer {key}", "user-agent": "diag/1.0"}) + out = _post( + FW_URL, body, {"authorization": f"Bearer {key}", "user-agent": "diag/1.0"} + ) choice = (out.get("choices") or [{}])[0] msg = choice.get("message") or {} text = msg.get("content") or "" @@ -55,9 +66,15 @@ def fw_complete(messages: list[dict], max_tokens: int = 32768) -> dict: "out": u.get("completion_tokens", 0), "cache_w": 0, "cache_r": 0, - "reasoning": (u.get("completion_tokens_details") or {}).get("reasoning_tokens", 0), + "reasoning": (u.get("completion_tokens_details") or {}).get( + "reasoning_tokens", 0 + ), } - stop = "max_tokens" if choice.get("finish_reason") == "length" else (choice.get("finish_reason") or "") + stop = ( + "max_tokens" + if choice.get("finish_reason") == "length" + else (choice.get("finish_reason") or "") + ) return {"text": text, "usage": usage, "stop": stop} @@ -73,7 +90,9 @@ def main() -> None: answers, usages, stops, qs_all = [], [], [], [] for batch, messages in batches: qa = cached( - "fireworks/kimi-k2p6", "qa-mono-prod", {"messages": messages, "effort": None}, + "fireworks/kimi-k2p6", + "qa-mono-prod", + {"messages": messages, "effort": None}, lambda m=messages: fw_complete(m), args.fresh, ) @@ -84,34 +103,57 @@ def main() -> None: rows = [ { - "model": "fireworks/kimi-k2p6", "cond": cond, "pos_rel": q["pos_rel"], "q": q["q"], "answer": a, - "golds": q["golds"], "em": squad.exact_match(a, q["golds"]), "f1": squad.f1(a, q["golds"]), + "model": "fireworks/kimi-k2p6", + "cond": cond, + "pos_rel": q["pos_rel"], + "q": q["q"], + "answer": a, + "golds": q["golds"], + "em": squad.exact_match(a, q["golds"]), + "f1": squad.f1(a, q["golds"]), "abstained": "unreadable" in a.lower(), } for q, a in zip(qs_all, answers) ] - u = {k: sum(x[k] for x in usages) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} + u = { + k: sum(x[k] for x in usages) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } cost = u["in"] / 1e6 * PRICE_IN + u["out"] / 1e6 * PRICE_OUT quart = [] for lo, hi in ((0, 0.25), (0.25, 0.5), (0.5, 0.75), (0.75, 1.01)): sel = [r["f1"] for r in rows if lo <= r["pos_rel"] < hi] quart.append(sum(sel) / len(sel) if sel else float("nan")) summary = { - "cond": cond, "n": len(rows), "imgs": len(pngs), + "cond": cond, + "n": len(rows), + "imgs": len(pngs), "em": sum(r["em"] for r in rows) / len(rows), "f1": sum(r["f1"] for r in rows) / len(rows), "abst": sum(r["abstained"] for r in rows), - "tok_in": u["in"], "tok_out": u["out"], "reas": u["reasoning"], "cost": cost, - "stop": next((s for s in stops if s == "max_tokens"), stops[-1] if stops else ""), - "q1": quart[0], "q2": quart[1], "q3": quart[2], "q4": quart[3], + "tok_in": u["in"], + "tok_out": u["out"], + "reas": u["reasoning"], + "cost": cost, + "stop": next( + (s for s in stops if s == "max_tokens"), stops[-1] if stops else "" + ), + "q1": quart[0], + "q2": quart[1], + "q3": quart[2], + "q4": quart[3], } out_dir = RESULTS / f"diag-kimi-fireworks-{args.shape}" out_dir.mkdir(parents=True, exist_ok=True) (out_dir / "records.jsonl").write_text("\n".join(json.dumps(r) for r in rows)) (out_dir / "summary.json").write_text(json.dumps([summary], indent=1)) - print(f"{cond} imgs={summary['imgs']} f1={summary['f1']:.3f} em={summary['em']:.3f} " - f"abst={summary['abst']} ${summary['cost']:.2f} stop={summary['stop']}") - print("F1 by quartile: " + " ".join(f"q{i + 1}={v:.3f}" for i, v in enumerate(quart))) + print( + f"{cond} imgs={summary['imgs']} f1={summary['f1']:.3f} em={summary['em']:.3f} " + f"abst={summary['abst']} ${summary['cost']:.2f} stop={summary['stop']}" + ) + print( + "F1 by quartile: " + " ".join(f"q{i + 1}={v:.3f}" for i, v in enumerate(quart)) + ) if __name__ == "__main__": diff --git a/packages/snapcompact/research/diag_kimi_probe.py b/packages/snapcompact/research/diag_kimi_probe.py index 4dbd960fb..bcd420caf 100644 --- a/packages/snapcompact/research/diag_kimi_probe.py +++ b/packages/snapcompact/research/diag_kimi_probe.py @@ -45,10 +45,17 @@ def chat(route: str, model: str, content: list[dict], max_tokens: int = 2048) -> if route == "openrouter": url, key = OPENROUTER_URL, load_env_key("OPENROUTER_API_KEY") elif route == "fireworks": - url, key = "https://api.fireworks.ai/inference/v1/chat/completions", load_env_key("FIREWORKS_API_KEY") + url, key = ( + "https://api.fireworks.ai/inference/v1/chat/completions", + load_env_key("FIREWORKS_API_KEY"), + ) else: url, key = f"{MOONSHOT_BASE}/chat/completions", load_env_key("KIMI_API_KEY") - body = {"model": model, "messages": [{"role": "user", "content": content}], "max_tokens": max_tokens} + body = { + "model": model, + "messages": [{"role": "user", "content": content}], + "max_tokens": max_tokens, + } if route == "openrouter" and PROVIDER: body["provider"] = {"order": [PROVIDER], "allow_fallbacks": False} out = _post(url, body, {"authorization": f"Bearer {key}", "user-agent": "diag/1.0"}) @@ -57,8 +64,13 @@ def chat(route: str, model: str, content: list[dict], max_tokens: int = 2048) -> text = msg.get("content") or "" if isinstance(text, list): text = "".join(p.get("text", "") for p in text if isinstance(p, dict)) - return {"text": text, "usage": out.get("usage", {}), "finish": choice.get("finish_reason"), - "provider": out.get("provider"), "model": out.get("model")} + return { + "text": text, + "usage": out.get("usage", {}), + "finish": choice.get("finish_reason"), + "provider": out.get("provider"), + "model": out.get("model"), + } def small_frames(px: int) -> list[Path]: @@ -75,28 +87,39 @@ def small_frames(px: int) -> list[Path]: outs.append(q) return outs + def img_block(p: Path) -> dict: - return {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{_png_b64(p)}"}} + return { + "type": "image_url", + "image_url": {"url": f"data:image/png;base64,{_png_b64(p)}"}, + } def cmd_models() -> None: import urllib.request key = load_env_key("KIMI_API_KEY") - req = urllib.request.Request(f"{MOONSHOT_BASE}/models", headers={"authorization": f"Bearer {key}"}) + req = urllib.request.Request( + f"{MOONSHOT_BASE}/models", headers={"authorization": f"Bearer {key}"} + ) with urllib.request.urlopen(req, timeout=60) as resp: out = json.loads(resp.read()) for m in out.get("data", []): print(m.get("id")) -def cmd_tokens(route: str, model: str, counts: list[int], px: int | None = None) -> None: +def cmd_tokens( + route: str, model: str, counts: list[int], px: int | None = None +) -> None: pngs = small_frames(px) if px else frames() for k in counts: content = [ {"type": "text", "text": f"This message has some images attached."}, *(img_block(p) for p in pngs[:k]), - {"type": "text", "text": "How many images are attached to this message? Reply with just the integer."}, + { + "type": "text", + "text": "How many images are attached to this message? Reply with just the integer.", + }, ] try: r = chat(route, model, content) @@ -104,18 +127,26 @@ def cmd_tokens(route: str, model: str, counts: list[int], px: int | None = None) print(f"K={k:>2} ERROR {e}") continue u = r["usage"] - print(f"K={k:>2} prompt_tokens={u.get('prompt_tokens')} completion={u.get('completion_tokens')} " - f"provider={r.get('provider')} finish={r['finish']} answer={r['text'].strip()[:80]!r}") + print( + f"K={k:>2} prompt_tokens={u.get('prompt_tokens')} completion={u.get('completion_tokens')} " + f"provider={r.get('provider')} finish={r['finish']} answer={r['text'].strip()[:80]!r}" + ) def cmd_lastline(route: str, model: str, counts: list[int]) -> None: pngs = frames() for k in counts: content = [ - {"type": "text", "text": f"The attached {k} images contain text rendered in a monospace pixel font."}, + { + "type": "text", + "text": f"The attached {k} images contain text rendered in a monospace pixel font.", + }, *(img_block(p) for p in pngs[:k]), - {"type": "text", "text": f"Transcribe the first 10 words on the FIRST text row of the LAST image (image {k}). " - "If you cannot read it, reply exactly UNREADABLE."}, + { + "type": "text", + "text": f"Transcribe the first 10 words on the FIRST text row of the LAST image (image {k}). " + "If you cannot read it, reply exactly UNREADABLE.", + }, ] try: r = chat(route, model, content, max_tokens=4096) @@ -123,25 +154,37 @@ def cmd_lastline(route: str, model: str, counts: list[int]) -> None: print(f"K={k:>2} ERROR {e}") continue u = r["usage"] - print(f"K={k:>2} prompt_tokens={u.get('prompt_tokens')} finish={r['finish']}\n -> {r['text'].strip()[:200]!r}") + print( + f"K={k:>2} prompt_tokens={u.get('prompt_tokens')} finish={r['finish']}\n -> {r['text'].strip()[:200]!r}" + ) def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("cmd", choices=["models", "tokens", "lastline"]) - ap.add_argument("--px", type=int, default=None, help="downscale frames to this square size first") - ap.add_argument("--route", default="openrouter", choices=["openrouter", "moonshot", "fireworks"]) + ap.add_argument( + "--px", + type=int, + default=None, + help="downscale frames to this square size first", + ) + ap.add_argument( + "--route", default="openrouter", choices=["openrouter", "moonshot", "fireworks"] + ) ap.add_argument("--model", default=None) ap.add_argument("--counts", default="1,4,8,9,12,21") ap.add_argument("--provider", default=None) args = ap.parse_args() global PROVIDER PROVIDER = args.provider - model = args.model or { - "openrouter": OR_MODEL, - "fireworks": "accounts/fireworks/models/kimi-k2p6", - "moonshot": "kimi-k2.6", - }[args.route] + model = ( + args.model + or { + "openrouter": OR_MODEL, + "fireworks": "accounts/fireworks/models/kimi-k2p6", + "moonshot": "kimi-k2.6", + }[args.route] + ) if args.cmd == "models": cmd_models() elif args.cmd == "tokens": diff --git a/packages/snapcompact/research/exp01_patchalign.py b/packages/snapcompact/research/exp01_patchalign.py index 29431073c..e8a9a95f8 100644 --- a/packages/snapcompact/research/exp01_patchalign.py +++ b/packages/snapcompact/research/exp01_patchalign.py @@ -46,10 +46,14 @@ MODELS = { } FONTS = { "7x14": FontCfg("7x14", "7x14", 7, 14), # aligned: pitch = 14 px patch - "8x16": FontCfg("8x16", "spleen-8x16", 8, 16), # aligned: native 16 px font (Spleen) + "8x16": FontCfg( + "8x16", "spleen-8x16", 8, 16 + ), # aligned: native 16 px font (Spleen) "8on16": FontCfg("8on16", "8x13", 8, 16), # aligned: 8x13 glyphs, pitch 16 cell "6on7x14": FontCfg("6on7x14", "6x12", 7, 14), # aligned: 6x12 glyphs, 7x14 cell - "7x13": FontCfg("7x13", "7x13", 7, 13), # control for 7x14 (same glyph budget, pitch 13) + "7x13": FontCfg( + "7x13", "7x13", 7, 13 + ), # control for 7x14 (same glyph budget, pitch 13) "8x13": FontCfg("8x13", "8x13", 8, 13), # control for 8x16/8on16 } SPLEEN_URL = "https://raw.githubusercontent.com/fcambus/spleen/master/spleen-8x16.bdf" @@ -99,10 +103,20 @@ def ensure_spleen() -> None: tmp.replace(path) -def run_cell_chunk(model: str, cond: str, size: int, start: int, end: int, ctx: dict) -> list[dict]: +def run_cell_chunk( + model: str, cond: str, size: int, start: int, end: int, ctx: dict +) -> list[dict]: """One (model, condition, size, chunk) unit: render carrier, QA, score.""" - args, flow, paras, offsets, keys = ctx["args"], ctx["flow"], ctx["paras"], ctx["offsets"], ctx["keys"] - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + args, flow, paras, offsets, keys = ( + ctx["args"], + ctx["flow"], + ctx["paras"], + ctx["offsets"], + ctx["keys"], + ) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] chunk_text = flow[start:end] @@ -126,11 +140,19 @@ def run_cell_chunk(model: str, cond: str, size: int, start: int, end: int, ctx: } ] qa = cached( - model, "exp01-qa", {"messages": messages, "size": size, "effort": args.effort}, + model, + "exp01-qa", + {"messages": messages, "size": size, "effort": args.effort}, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, model, messages, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + model, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -164,8 +186,13 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out return { "n": n, @@ -190,7 +217,11 @@ def main() -> None: ap.add_argument("--max-tokens", type=int, default=32768) ap.add_argument("--effort", default=None) ap.add_argument("--fresh", action="store_true") - ap.add_argument("--report", action="store_true", help="reprint from cache only (re-runs cells; all should hit cache)") + ap.add_argument( + "--report", + action="store_true", + help="reprint from cache only (re-runs cells; all should hit cache)", + ) ap.add_argument("--env", default="~/.env") args = ap.parse_args() @@ -213,18 +244,32 @@ def main() -> None: for length in lengths: paras = all_paras[:length] flow, offsets = squad.build_flow(paras) - ctx = {"args": args, "flow": flow, "paras": paras, "offsets": offsets, "keys": keys, "length": length} + ctx = { + "args": args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + } for model in models: for cond, size in GRID: budget = capacity(FONTS[parse_img_condition(cond)[0]], size)[2] for start in range(0, len(flow), budget): - tasks.append((model, cond, size, start, min(start + budget, len(flow)), ctx)) - print(f"grid: {len(models)} models x {len(lengths)} lengths x {len(GRID)} cells = {len(tasks)} chunk tasks") + tasks.append( + (model, cond, size, start, min(start + budget, len(flow)), ctx) + ) + print( + f"grid: {len(models)} models x {len(lengths)} lengths x {len(GRID)} cells = {len(tasks)} chunk tasks" + ) records: list[dict] = [] done = 0 with ThreadPoolExecutor(args.workers) as pool: - futures = [pool.submit(run_cell_chunk, m, c, sz, s, e, ctx) for m, c, sz, s, e, ctx in tasks] + futures = [ + pool.submit(run_cell_chunk, m, c, sz, s, e, ctx) + for m, c, sz, s, e, ctx in tasks + ] for fut in futures: records.extend(fut.result()) done += 1 @@ -240,8 +285,12 @@ def main() -> None: for length in lengths: for cond, size in GRID: sub = [ - r for r in records - if r["model"] == model and r["length"] == length and r["cond"] == cond and r["size"] == size + r + for r in records + if r["model"] == model + and r["length"] == length + and r["cond"] == cond + and r["size"] == size ] if not sub: continue @@ -253,7 +302,9 @@ def main() -> None: **aggregate(sub, *MODELS[model]), } ) - (out_dir / "summary.json").write_text(json.dumps({"args": vars(args), "cells": cells}, indent=1)) + (out_dir / "summary.json").write_text( + json.dumps({"args": vars(args), "cells": cells}, indent=1) + ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) writer.writeheader() @@ -268,7 +319,13 @@ def main() -> None: row = f"{label:<22}" for model in models: cell = next( - (c for c in cells if c["model"] == model and c["length"] == length and c["condition"] == label), + ( + c + for c in cells + if c["model"] == model + and c["length"] == length + and c["condition"] == label + ), None, ) row += ( diff --git a/packages/snapcompact/research/exp02_surprisal.py b/packages/snapcompact/research/exp02_surprisal.py index 63fe857f6..b65ea093b 100644 --- a/packages/snapcompact/research/exp02_surprisal.py +++ b/packages/snapcompact/research/exp02_surprisal.py @@ -129,7 +129,8 @@ _GRAYS = [(185, 185, 185), (135, 135, 135), (75, 75, 75), (0, 0, 0)] _LIGHT = [0.72, 0.55, 0.40, 0.22] # lightness per bucket for sent hues _HUES = [0.0, 0.08, 0.3, 0.5, 0.62, 0.78] _SENT_SURP = [ - [tuple(int(c * 255) for c in colorsys.hls_to_rgb(h, l, 0.95)) for l in _LIGHT] for h in _HUES + [tuple(int(c * 255) for c in colorsys.hls_to_rgb(h, l, 0.95)) for l in _LIGHT] + for h in _HUES ] _WHITE = (255, 255, 255) @@ -158,7 +159,9 @@ def _shade_indices(text: str) -> list[int]: return out -def render_surp(text: str, cfg, cache: Path, size: int = 1568, sent_hues: bool = False) -> Image.Image: +def render_surp( + text: str, cfg, cache: Path, size: int = 1568, sent_hues: bool = False +) -> Image.Image: """bdf.render() with the boolean dim_mask generalized to surprisal buckets.""" glyphs, font_ascent = parse_bdf(ensure_font(cfg, cache)) ascent = cfg.ascent if cfg.ascent is not None else font_ascent @@ -257,9 +260,13 @@ def build_png(cond: str, render_text: str, size: int) -> Path: return png -def run_chunk(model: str, cond: str, start: int, end: int, png: Path, render_chars: int, ctx: dict) -> list[dict]: +def run_chunk( + model: str, cond: str, start: int, end: int, png: Path, render_chars: int, ctx: dict +) -> list[dict]: args, keys = ctx["args"], ctx["keys"] - questions = squad.sample_chunk_questions(ctx["paras"], ctx["offsets"], start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + ctx["paras"], ctx["offsets"], start, end, args.qpc, args.seed + ) if not questions: return [] q_block = "\n".join(f"{i + 1}. {q['q']}" for i, q in enumerate(questions)) @@ -269,18 +276,30 @@ def run_chunk(model: str, cond: str, start: int, end: int, png: Path, render_cha { "role": "user", "content": [ - {"text": load_prompt("exp02-qa-image.md").format(cols=cols, rows=rows, extra=extra)}, + { + "text": load_prompt("exp02-qa-image.md").format( + cols=cols, rows=rows, extra=extra + ) + }, {"image_path": png}, {"text": q_block}, ], } ] qa = cached( - model, "exp02-qa", {"messages": messages, "effort": args.effort}, + model, + "exp02-qa", + {"messages": messages, "effort": args.effort}, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, model, messages, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + model, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -315,8 +334,13 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out pages = sum(1 for r in records if "chunk_orig_chars" in r) orig_chars = sum(r.get("chunk_orig_chars", 0) for r in records) @@ -371,13 +395,21 @@ def main() -> None: for length in lengths: paras = all_paras[:length] flow, offsets = squad.build_flow(paras) - ctx = {"args": args, "paras": paras, "offsets": offsets, "keys": keys, "length": length} + ctx = { + "args": args, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + } for cond in conditions: chunks = plan_chunks(flow, cond, args.size) for start, end in chunks: orig = flow[start:end] render_text = transform(orig) if cond == "img-6x10-disemv" else orig - png = build_png(cond, render_text, args.size) # pre-render: no tmp races in pool + png = build_png( + cond, render_text, args.size + ) # pre-render: no tmp races in pool if cond == "img-6x10-disemv": print( f" len={length} disemv chunk [{start},{end}): {end - start} orig -> " @@ -385,7 +417,9 @@ def main() -> None: ) for model in models: tasks.append((model, cond, start, end, png, len(render_text), ctx)) - print(f"grid: {len(models)} models x {len(lengths)} lengths x {len(conditions)} conditions = {len(tasks)} chunk tasks") + print( + f"grid: {len(models)} models x {len(lengths)} lengths x {len(conditions)} conditions = {len(tasks)} chunk tasks" + ) records: list[dict] = [] with ThreadPoolExecutor(args.workers) as pool: @@ -402,17 +436,34 @@ def main() -> None: for model in models: for length in lengths: for cond in conditions: - sub = [r for r in records if r["model"] == model and r["length"] == length and r["cond"] == cond] + sub = [ + r + for r in records + if r["model"] == model + and r["length"] == length + and r["cond"] == cond + ] if not sub: continue - cells.append({"model": model, "length": length, "condition": cond, **aggregate(sub, *MODELS[model])}) - (out_dir / "summary.json").write_text(json.dumps({"args": vars(args), "cells": cells}, indent=1)) + cells.append( + { + "model": model, + "length": length, + "condition": cond, + **aggregate(sub, *MODELS[model]), + } + ) + (out_dir / "summary.json").write_text( + json.dumps({"args": vars(args), "cells": cells}, indent=1) + ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) writer.writeheader() writer.writerows(cells) - print(f"\n{'model':<26}{'len':>5}{'condition':<22}{'n':>4}{'EM':>7}{'F1':>7}{'se':>7}{'cost$':>8}{'dF1 vs base':>13}") + print( + f"\n{'model':<26}{'len':>5}{'condition':<22}{'n':>4}{'EM':>7}{'F1':>7}{'se':>7}{'cost$':>8}{'dF1 vs base':>13}" + ) for c in cells: base = BASELINE.get((c["model"], c["length"])) d = f"{c['f1'] - base[0]:+.3f}" if base else "-" diff --git a/packages/snapcompact/research/exp03_numhard.py b/packages/snapcompact/research/exp03_numhard.py index 115fbb96c..e1afc5578 100644 --- a/packages/snapcompact/research/exp03_numhard.py +++ b/packages/snapcompact/research/exp03_numhard.py @@ -63,12 +63,16 @@ def number_mask(text: str) -> list[bool]: if ch.isdigit(): mask[i] = True elif ch in _NUM_PUNCT: - if (i > 0 and text[i - 1].isdigit()) or (i + 1 < len(text) and text[i + 1].isdigit()): + if (i > 0 and text[i - 1].isdigit()) or ( + i + 1 < len(text) and text[i + 1].isdigit() + ): mask[i] = True return mask -def render_hard(text: str, cfg, cache: Path, size: int, variant: str, hard: str) -> Image.Image: +def render_hard( + text: str, cfg, cache: Path, size: int, variant: str, hard: str +) -> Image.Image: """Copy of bdf.render() restricted to white-bg variants (sent/bw), with a digit-hardening pass: `numbold` double-strikes masked glyphs, `numred` recolors them pure red.""" @@ -132,7 +136,9 @@ def chunk_png(chunk_text: str, size: int, base: str, hard: str) -> Path: def run_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> list[dict]: args, flow = ctx["args"], ctx["flow"] - questions = squad.sample_chunk_questions(ctx["paras"], ctx["offsets"], start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + ctx["paras"], ctx["offsets"], start, end, args.qpc, args.seed + ) if not questions: return [] chunk_text = flow[start:end] @@ -151,11 +157,19 @@ def run_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> list[di } ] qa = cached( - model, f"{EXP}-qa", {"messages": messages, "effort": args.effort}, + model, + f"{EXP}-qa", + {"messages": messages, "effort": args.effort}, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(ctx["keys"], model, messages, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + ctx["keys"], + model, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -209,19 +223,29 @@ def load_baseline(model: str, lengths: list[int]) -> list[dict]: return out -def numeric_subset_cells(records: list[dict], models: list[str], lengths: list[int], conditions: list[str]) -> list[dict]: +def numeric_subset_cells( + records: list[dict], models: list[str], lengths: list[int], conditions: list[str] +) -> list[dict]: """Per (model, length): baseline vs each condition, restricted to numeric-gold questions present in BOTH runs (matched by question text).""" cells = [] for model in models: base = load_baseline(model, lengths) for length in lengths: - base_num = {r["q"]: r for r in base if r["length"] == length and is_numeric_gold(r["golds"])} + base_num = { + r["q"]: r + for r in base + if r["length"] == length and is_numeric_gold(r["golds"]) + } for cond in conditions: mine = [ - r for r in records - if r["model"] == model and r["length"] == length and r["cond"] == cond - and is_numeric_gold(r["golds"]) and r["q"] in base_num + r + for r in records + if r["model"] == model + and r["length"] == length + and r["cond"] == cond + and is_numeric_gold(r["golds"]) + and r["q"] in base_num ] if not mine: continue @@ -239,7 +263,9 @@ def numeric_subset_cells(records: list[dict], models: list[str], lengths: list[i return cells -def save_sample(records_ctx_flow: str, size: int, base: str, hard: str, out_dir: Path) -> Path: +def save_sample( + records_ctx_flow: str, size: int, base: str, hard: str, out_dir: Path +) -> Path: """Crop a digit-dense region from the first chunk's PNG, 4x nearest upscale.""" cols, rows, cap = capacity(FONT, size) text = records_ctx_flow[:cap] @@ -257,7 +283,9 @@ def save_sample(records_ctx_flow: str, size: int, base: str, hard: str, out_dir: img = Image.open(png) x0 = max(0, min(col, cols - win) * FONT.adv) y0 = max(0, (row - 1) * FONT.pitch) - crop = img.crop((x0, y0, min(x0 + win * FONT.adv, size), min(y0 + 4 * FONT.pitch, size))) + crop = img.crop( + (x0, y0, min(x0 + win * FONT.adv, size), min(y0 + 4 * FONT.pitch, size)) + ) crop = crop.resize((crop.width * 4, crop.height * 4), Image.NEAREST) out = out_dir / f"sample-{base}-{hard}.png" crop.save(out) @@ -302,12 +330,23 @@ def main() -> None: paras = all_paras[:length] flow, offsets = squad.build_flow(paras) flows[length] = flow - ctx = {"args": args, "flow": flow, "paras": paras, "offsets": offsets, "keys": keys, "length": length} + ctx = { + "args": args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + } for model in models: for cond in conditions: for start in range(0, len(flow), budget): - tasks.append((model, cond, start, min(start + budget, len(flow)), ctx)) - print(f"{EXP}: {len(models)} models x {len(lengths)} lengths x {len(conditions)} conditions = {len(tasks)} chunk tasks") + tasks.append( + (model, cond, start, min(start + budget, len(flow)), ctx) + ) + print( + f"{EXP}: {len(models)} models x {len(lengths)} lengths x {len(conditions)} conditions = {len(tasks)} chunk tasks" + ) records: list[dict] = [] with ThreadPoolExecutor(args.workers) as pool: @@ -325,9 +364,22 @@ def main() -> None: for model in models: for length in lengths: for cond in conditions: - sub = [r for r in records if r["model"] == model and r["length"] == length and r["cond"] == cond] + sub = [ + r + for r in records + if r["model"] == model + and r["length"] == length + and r["cond"] == cond + ] if sub: - cells.append({"model": model, "length": length, "condition": cond, **aggregate(sub, *MODELS[model])}) + cells.append( + { + "model": model, + "length": length, + "condition": cond, + **aggregate(sub, *MODELS[model]), + } + ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) writer.writeheader() @@ -349,7 +401,9 @@ def main() -> None: "args": vars(args), "cells": cells, "numeric_subset": num_cells, - "baseline_overall": {f"{m}|{l}": v for (m, l), v in base_overall.items()}, + "baseline_overall": { + f"{m}|{l}": v for (m, l), v in base_overall.items() + }, }, indent=1, ) @@ -358,10 +412,14 @@ def main() -> None: samples = [] for cond in conditions: base, hard = parse_cond(cond) - samples.append(str(save_sample(flows[lengths[0]], args.size, base, hard, out_dir))) + samples.append( + str(save_sample(flows[lengths[0]], args.size, base, hard, out_dir)) + ) print("\n== overall ==") - print(f"{'model':<24}{'len':>5}{'condition':<28}{'n':>4}{'EM':>7}{'F1':>7}{'se':>7}{'cost$':>8}{'base F1':>9}{'d':>7}") + print( + f"{'model':<24}{'len':>5}{'condition':<28}{'n':>4}{'EM':>7}{'F1':>7}{'se':>7}{'cost$':>8}{'base F1':>9}{'d':>7}" + ) for c in cells: b = base_overall[(c["model"], c["length"])] print( @@ -369,7 +427,9 @@ def main() -> None: f"{c['f1_se']:>7.3f}{c['cost_usd']:>8.3f}{b['f1']:>9.3f}{c['f1'] - b['f1']:>+7.3f}" ) print("\n== numeric-gold subset (matched questions vs img-6x10-sent baseline) ==") - print(f"{'model':<24}{'len':>5}{'condition':<28}{'n':>4}{'F1':>7}{'se':>7}{'base F1':>9}{'base se':>8}{'d':>7}") + print( + f"{'model':<24}{'len':>5}{'condition':<28}{'n':>4}{'F1':>7}{'se':>7}{'base F1':>9}{'base se':>8}{'d':>7}" + ) for c in num_cells: print( f"{c['model']:<24}{c['length']:>5}{c['condition']:<28}{c['n']:>4}{c['f1']:>7.3f}{c['f1_se']:>7.3f}" diff --git a/packages/snapcompact/research/exp04_layout.py b/packages/snapcompact/research/exp04_layout.py index 0d8a1d861..c89d382c3 100644 --- a/packages/snapcompact/research/exp04_layout.py +++ b/packages/snapcompact/research/exp04_layout.py @@ -170,12 +170,16 @@ def run_page(model: str, cond: str, page: tuple[int, int], ctx: dict) -> list[di i, j = page start = offsets[i] end = offsets[j - 1] + len(paras[j - 1]["ctx"]) - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] variant = cond.removeprefix("img-6x10-") lines = ctx["lines"][page] - page_key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size)) + page_key = sha8( + cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size) + ) png = CACHE / f"exp04-{variant}-{page_key}.png" if not png.exists() or png.stat().st_size == 0: tmp = png.with_suffix(".tmp.png") @@ -188,18 +192,30 @@ def run_page(model: str, cond: str, page: tuple[int, int], ctx: dict) -> list[di { "role": "user", "content": [ - {"text": load_prompt("exp04-qa-image.md").format(col_w=col_w, rows=rows)}, + { + "text": load_prompt("exp04-qa-image.md").format( + col_w=col_w, rows=rows + ) + }, {"image_path": png}, {"text": q_block}, ], } ] qa = cached( - model, "exp04-qa", {"messages": messages, "effort": args.effort}, + model, + "exp04-qa", + {"messages": messages, "effort": args.effort}, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, model, messages, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + model, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -232,8 +248,13 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out return { "n": n, @@ -260,7 +281,11 @@ def main() -> None: ap.add_argument("--max-tokens", type=int, default=32768) ap.add_argument("--effort", default=None) ap.add_argument("--fresh", action="store_true") - ap.add_argument("--render-only", action="store_true", help="render pages + capacity stats, no API") + ap.add_argument( + "--render-only", + action="store_true", + help="render pages + capacity stats, no API", + ) ap.add_argument("--env", default="~/.env") args = ap.parse_args() @@ -290,7 +315,9 @@ def main() -> None: flow, offsets = squad.build_flow(paras) pages = pack_pages(paras, col_w, max_lines) page_lines = {pg: layout_page(paras[pg[0] : pg[1]], col_w) for pg in pages} - page_chars = [offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] for i, j in pages] + page_chars = [ + offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] for i, j in pages + ] capacity_stats[length] = { "pages": len(pages), "chars_per_page": page_chars, @@ -300,15 +327,21 @@ def main() -> None: "grid_pages": -(-len(flow) // grid_cap), } ctx = { - "args": args, "paras": paras, "offsets": offsets, "keys": keys, - "length": length, "lines": page_lines, + "args": args, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + "lines": page_lines, } for model in models: for cond in conditions: for pg in pages: tasks.append((model, cond, pg, ctx)) - print(f"layout: {cols} cols -> 2 x {col_w} + gutter {GUTTER}; {max_lines} line slots/page") + print( + f"layout: {cols} cols -> 2 x {col_w} + gutter {GUTTER}; {max_lines} line slots/page" + ) for length, st in capacity_stats.items(): print( f" len {length}: {st['pages']} doc pages (mean {st['mean_chars_page']} chars/page; " @@ -322,7 +355,11 @@ def main() -> None: variant = cond.removeprefix("img-6x10-") i, j = pages[0] lines = layout_page(paras[i:j], col_w) - key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size)) + key = sha8( + cond, + json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), + str(args.size), + ) png = CACHE / f"exp04-{variant}-{key}.png" tmp = png.with_suffix(".tmp.png") render_doc(lines, args.size, variant, CACHE).save(tmp) @@ -348,12 +385,27 @@ def main() -> None: for model in models: for length in lengths: for cond in conditions: - sub = [r for r in records if r["model"] == model and r["length"] == length and r["cond"] == cond] + sub = [ + r + for r in records + if r["model"] == model + and r["length"] == length + and r["cond"] == cond + ] if not sub: continue - cells.append({"model": model, "length": length, "condition": cond, **aggregate(sub, *MODELS[model])}) + cells.append( + { + "model": model, + "length": length, + "condition": cond, + **aggregate(sub, *MODELS[model]), + } + ) (out_dir / "summary.json").write_text( - json.dumps({"args": vars(args), "capacity": capacity_stats, "cells": cells}, indent=1) + json.dumps( + {"args": vars(args), "capacity": capacity_stats, "cells": cells}, indent=1 + ) ) import csv diff --git a/packages/snapcompact/research/exp05_anchors.py b/packages/snapcompact/research/exp05_anchors.py index 2cea6a75f..61d984632 100644 --- a/packages/snapcompact/research/exp05_anchors.py +++ b/packages/snapcompact/research/exp05_anchors.py @@ -35,7 +35,17 @@ HERE = Path(__file__).resolve().parent sys.path.insert(0, str(HERE)) import squad # noqa: E402 -from bdf import _BLACK, _DARK, _WHITE, FontCfg, _row_palette, _sentence_indices, ensure_font, parse_bdf, render # noqa: E402 +from bdf import ( + _BLACK, + _DARK, + _WHITE, + FontCfg, + _row_palette, + _sentence_indices, + ensure_font, + parse_bdf, + render, +) # noqa: E402 from final import MODELS, cached # noqa: E402 from providers import llm_complete, load_env_key # noqa: E402 from run import CACHE, FONTS, QA_CACHE, RESULTS, load_prompt, sha8 # noqa: E402 @@ -63,7 +73,9 @@ def ruler_capacity(cfg: FontCfg, size: int = SIZE) -> tuple[int, int, int]: return cols, rows, cols * rows -def render_ruler(text: str, cfg: FontCfg, cache: Path, size: int = SIZE, variant: str = "sent") -> Image.Image: +def render_ruler( + text: str, cfg: FontCfg, cache: Path, size: int = SIZE, variant: str = "sent" +) -> Image.Image: """bdf.render() with a left row-number ruler every RULER_STEP rows. Content glyphs are shifted right by MARGIN_COLS cells; row indices are @@ -73,9 +85,15 @@ def render_ruler(text: str, cfg: FontCfg, cache: Path, size: int = SIZE, variant ascent = cfg.ascent if cfg.ascent is not None else font_ascent cols, rows, cap = ruler_capacity(cfg, size) text = text[:cap] - sent_idx = _sentence_indices(text) if variant in ("sent", "dark-sent", "sent-dim") else None + sent_idx = ( + _sentence_indices(text) + if variant in ("sent", "dark-sent", "sent-dim") + else None + ) sent_palette = _DARK - img = Image.new("RGB", (size, size), _BLACK if variant in ("dark", "dark-sent") else _WHITE) + img = Image.new( + "RGB", (size, size), _BLACK if variant in ("dark", "dark-sent") else _WHITE + ) px = img.load() def draw_glyph(ch: str, cell_col: int, y0: int, fg: tuple[int, int, int]) -> None: @@ -160,9 +178,19 @@ def _ensure_png(png: Path, make) -> None: tmp.replace(png) -def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> list[dict]: - args, flow, paras, offsets, keys = ctx["args"], ctx["flow"], ctx["paras"], ctx["offsets"], ctx["keys"] - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) +def run_cell_chunk( + model: str, cond: str, start: int, end: int, ctx: dict +) -> list[dict]: + args, flow, paras, offsets, keys = ( + ctx["args"], + ctx["flow"], + ctx["paras"], + ctx["offsets"], + ctx["keys"], + ) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] chunk_text = flow[start:end] @@ -171,9 +199,13 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li if cond == COND_RULER: cols, rows, _ = ruler_capacity(FONT, args.size) png = CACHE / f"exp05-ruler-{sha8(chunk_text, str(args.size))}.png" - _ensure_png(png, lambda: render_ruler(chunk_text, FONT, CACHE, args.size, "sent")) + _ensure_png( + png, lambda: render_ruler(chunk_text, FONT, CACHE, args.size, "sent") + ) last_label = (rows - 1) // RULER_STEP * RULER_STEP - prompt = load_prompt("exp05-qa-image.md").format(cols=cols, rows=rows, last_label=last_label) + prompt = load_prompt("exp05-qa-image.md").format( + cols=cols, rows=rows, last_label=last_label + ) else: # control: baseline render, anti-transcription prompt only cols, rows = args.size // FONT.adv, args.size // FONT.pitch png = CACHE / f"exp05-ctl-{sha8(chunk_text, str(args.size))}.png" @@ -187,11 +219,19 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li } ] qa = cached( - model, f"exp05-qa-{cond}", {"messages": messages, "effort": args.effort}, + model, + f"exp05-qa-{cond}", + {"messages": messages, "effort": args.effort}, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, model, messages, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + model, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -199,7 +239,9 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li answers, claimed_rows = parse_answers_rows(qa["text"], len(questions)) records = [] for q, a, crow in zip(questions, answers, claimed_rows): - trow = gold_row(chunk_text, q, end - start, cols) if cond == COND_RULER else None + trow = ( + gold_row(chunk_text, q, end - start, cols) if cond == COND_RULER else None + ) records.append( { "model": model, @@ -227,16 +269,27 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out - loc = [(r["claimed_row"], r["true_row"]) for r in records if r["claimed_row"] is not None and r["true_row"] is not None] + loc = [ + (r["claimed_row"], r["true_row"]) + for r in records + if r["claimed_row"] is not None and r["true_row"] is not None + ] row_stats = {} if loc: errs = [abs(c - t) for c, t in loc] row_stats = { "row_n": len(loc), - "row_claimed_frac": round(sum(r["claimed_row"] is not None for r in records) / n, 3), + "row_claimed_frac": round( + sum(r["claimed_row"] is not None for r in records) / n, 3 + ), "row_mae": round(sum(errs) / len(errs), 2), "row_within2": round(sum(e <= 2 for e in errs) / len(errs), 3), "row_within5": round(sum(e <= 5 for e in errs) / len(errs), 3), @@ -291,16 +344,27 @@ def main() -> None: for length in lengths: paras = all_paras[:length] flow, offsets = squad.build_flow(paras) - ctx = {"args": args, "flow": flow, "paras": paras, "offsets": offsets, "keys": keys, "length": length} + ctx = { + "args": args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + } for model in models: for cond in conditions: for start in range(0, len(flow), budget): - tasks.append((model, cond, start, min(start + budget, len(flow)), ctx)) + tasks.append( + (model, cond, start, min(start + budget, len(flow)), ctx) + ) print(f"grid: {len(tasks)} chunk tasks, chunk budget {budget} chars") records: list[dict] = [] with ThreadPoolExecutor(args.workers) as pool: - futures = [pool.submit(run_cell_chunk, m, c, s, e, ctx) for m, c, s, e, ctx in tasks] + futures = [ + pool.submit(run_cell_chunk, m, c, s, e, ctx) for m, c, s, e, ctx in tasks + ] for i, fut in enumerate(futures): records.extend(fut.result()) print(f" {i + 1}/{len(tasks)} tasks", flush=True) @@ -313,10 +377,21 @@ def main() -> None: for model in models: for length in lengths: for cond in conditions: - sub = [r for r in records if r["model"] == model and r["length"] == length and r["cond"] == cond] + sub = [ + r + for r in records + if r["model"] == model + and r["length"] == length + and r["cond"] == cond + ] if not sub: continue - cell = {"model": model, "length": length, "condition": cond, **aggregate(sub, *MODELS[model])} + cell = { + "model": model, + "length": length, + "condition": cond, + **aggregate(sub, *MODELS[model]), + } base = BASELINE.get((model, length)) if base: cell["base_f1"] = base[0] @@ -324,8 +399,13 @@ def main() -> None: cell["base_cost"] = base[2] cell["d_cost"] = round(cell["cost_usd"] - base[2], 4) cells.append(cell) - (out_dir / "summary.json").write_text(json.dumps({"args": vars(args), "cells": cells}, indent=1)) - fieldnames = sorted({k for c in cells for k in c}, key=lambda k: (k not in ("model", "length", "condition"), k)) + (out_dir / "summary.json").write_text( + json.dumps({"args": vars(args), "cells": cells}, indent=1) + ) + fieldnames = sorted( + {k for c in cells for k in c}, + key=lambda k: (k not in ("model", "length", "condition"), k), + ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=fieldnames) writer.writeheader() @@ -335,7 +415,11 @@ def main() -> None: print( f"{c['model']:>26} L{c['length']:<4}{c['condition']:<24} n={c['n']:<4} em={c['em']:.3f} f1={c['f1']:.3f}±{c['f1_se']:.3f}" f" cost=${c['cost_usd']:.3f} out={c['tok_out']} reas={c['tok_reasoning']}" - + (f" rowMAE={c['row_mae']} w5={c['row_within5']}" if "row_mae" in c else "") + + ( + f" rowMAE={c['row_mae']} w5={c['row_within5']}" + if "row_mae" in c + else "" + ) ) print(f"\n-> {out_dir}/records.jsonl, matrix.csv, summary.json") diff --git a/packages/snapcompact/research/exp06_rolecolor.py b/packages/snapcompact/research/exp06_rolecolor.py index e1c81be2c..27b889ef6 100644 --- a/packages/snapcompact/research/exp06_rolecolor.py +++ b/packages/snapcompact/research/exp06_rolecolor.py @@ -60,7 +60,10 @@ FONT = FONTS["6x10"] ROLES = ("user", "assistant", "tool") TAGS = {"user": "user", "assistant": "asst", "tool": "tool"} # all "[xxxx] " = 7 chars ROLE_HUES = {"user": 0.62, "assistant": 0.33, "tool": 0.02} -ROLE_RGB = {r: tuple(int(c * 255) for c in colorsys.hls_to_rgb(h, 0.27, 0.90)) for r, h in ROLE_HUES.items()} +ROLE_RGB = { + r: tuple(int(c * 255) for c in colorsys.hls_to_rgb(h, 0.27, 0.90)) + for r, h in ROLE_HUES.items() +} _WHITE = (255, 255, 255) ENCODING = { @@ -76,7 +79,10 @@ ENCODING = { "no role information. Use your best guess." ), } -QA_PROMPT = {"img-6x10-role": "exp06-qa-image.md", "img-6x10-tagbw": "exp06-qa-image-tag.md"} +QA_PROMPT = { + "img-6x10-role": "exp06-qa-image.md", + "img-6x10-tagbw": "exp06-qa-image-tag.md", +} def assign_roles(n: int, seed: int) -> list[str]: @@ -104,10 +110,16 @@ def build_chunks(paras: list[dict], budget: int) -> list[tuple[int, int]]: return chunks -def sample_questions(paras: list[dict], offsets: list[int], start: int, end: int, n: int, seed: int) -> list[dict]: +def sample_questions( + paras: list[dict], offsets: list[int], start: int, end: int, n: int, seed: int +) -> list[dict]: """squad.sample_chunk_questions with the source passage index recorded (same rng sequence).""" rng = random.Random(seed * 1_000_003 + start) - eligible = [i for i in range(len(offsets)) if offsets[i] >= start and offsets[i] + len(paras[i]["ctx"]) <= end] + eligible = [ + i + for i in range(len(offsets)) + if offsets[i] >= start and offsets[i] + len(paras[i]["ctx"]) <= end + ] if not eligible: return [] n = min(n, len(eligible)) @@ -127,7 +139,9 @@ def sample_questions(paras: list[dict], offsets: list[int], start: int, end: int return picked -def render_role(text: str, colors: list[tuple[int, int, int]], size: int) -> Image.Image: +def render_role( + text: str, colors: list[tuple[int, int, int]], size: int +) -> Image.Image: """bdf.render() copy, simplified: white background, per-character glyph color.""" glyphs, font_ascent = parse_bdf(ensure_font(FONT, CACHE)) ascent = FONT.ascent if FONT.ascent is not None else font_ascent @@ -168,7 +182,11 @@ def chunk_carriers(paras: list[dict], roles: list[str], a: int, b: int) -> dict: plain_parts.append(seg) colors.extend([ROLE_RGB[roles[i]]] * len(seg)) tagged_parts.append(f"[{TAGS[roles[i]]}] {seg}") - return {"plain": "".join(plain_parts), "colors": colors, "tagged": "".join(tagged_parts)} + return { + "plain": "".join(plain_parts), + "colors": colors, + "tagged": "".join(tagged_parts), + } def atomic_png(png: Path, make) -> Path: @@ -201,7 +219,9 @@ def norm_role(answer: str) -> str: return a -def run_cell(model: str, cond: str, length: int, ci: int, chunk: dict, args, keys) -> list[dict]: +def run_cell( + model: str, cond: str, length: int, ci: int, chunk: dict, args, keys +) -> list[dict]: """One (model, cond, chunk): content QA (role/tagbw only) + provenance QA.""" questions, car = chunk["questions"], chunk["car"] if not questions: @@ -224,9 +244,14 @@ def run_cell(model: str, cond: str, length: int, ci: int, chunk: dict, args, key } ] qa = cached( - model, "exp06-qa", {"cond": cond, "messages": messages}, + model, + "exp06-qa", + {"cond": cond, "messages": messages}, lambda: dict( - zip(("text", "usage", "stop"), llm_complete(keys, model, messages, max_tokens=args.max_tokens)) + zip( + ("text", "usage", "stop"), + llm_complete(keys, model, messages, max_tokens=args.max_tokens), + ) ), args.fresh, ) @@ -237,16 +262,25 @@ def run_cell(model: str, cond: str, length: int, ci: int, chunk: dict, args, key { "role": "user", "content": [ - {"text": load_prompt("exp06-prov-image.md").format(cols=cols, rows=rows, encoding=ENCODING[cond])}, + { + "text": load_prompt("exp06-prov-image.md").format( + cols=cols, rows=rows, encoding=ENCODING[cond] + ) + }, {"image_path": png}, {"text": q_block}, ], } ] prov = cached( - model, "exp06-prov", {"cond": cond, "messages": prov_messages}, + model, + "exp06-prov", + {"cond": cond, "messages": prov_messages}, lambda: dict( - zip(("text", "usage", "stop"), llm_complete(keys, model, prov_messages, max_tokens=args.max_tokens)) + zip( + ("text", "usage", "stop"), + llm_complete(keys, model, prov_messages, max_tokens=args.max_tokens), + ) ), args.fresh, ) @@ -279,10 +313,17 @@ def run_cell(model: str, cond: str, length: int, ci: int, chunk: dict, args, key return records -def phase_cost(records: list[dict], phase: str, price_in: float, price_out: float) -> tuple[dict, float]: +def phase_cost( + records: list[dict], phase: str, price_in: float, price_out: float +) -> tuple[dict, float]: us = [u for r in records if "usage" in r for u in r["usage"] if u["phase"] == phase] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok["out"] / 1e6 * price_out + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost = ( + tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"] + ) / 1e6 * price_in + tok["out"] / 1e6 * price_out return tok, cost @@ -291,7 +332,11 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: f1s = [r["f1"] for r in records if r["f1"] is not None] if f1s: mean_f1 = sum(f1s) / len(f1s) - se = (sum((x - mean_f1) ** 2 for x in f1s) / (len(f1s) * (len(f1s) - 1))) ** 0.5 if len(f1s) > 1 else 0.0 + se = ( + (sum((x - mean_f1) ** 2 for x in f1s) / (len(f1s) * (len(f1s) - 1))) ** 0.5 + if len(f1s) > 1 + else 0.0 + ) em = sum(r["em"] for r in records if r["em"] is not None) / len(f1s) abstained = sum(r["abstained"] for r in records if r["abstained"] is not None) else: @@ -353,25 +398,35 @@ def main() -> None: for ci, (a, b) in enumerate(build_chunks(paras, budget)): start, end = offsets[a], offsets[b - 1] + len(paras[b - 1]["ctx"]) chunk = { - "questions": sample_questions(paras, offsets, start, end, args.qpc, args.seed), + "questions": sample_questions( + paras, offsets, start, end, args.qpc, args.seed + ), "car": chunk_carriers(paras, roles, a, b), "roles": roles, } for model in models: for cond in conditions: tasks.append((model, cond, length, ci, chunk)) - print(f"grid: {len(models)} models x {len(lengths)} lengths x {len(conditions)} conditions = {len(tasks)} cells") + print( + f"grid: {len(models)} models x {len(lengths)} lengths x {len(conditions)} conditions = {len(tasks)} cells" + ) records: list[dict] = [] failed = 0 with ThreadPoolExecutor(args.workers) as pool: - futures = [pool.submit(run_cell, m, c, ln, ci, ch, args, keys) for m, c, ln, ci, ch in tasks] + futures = [ + pool.submit(run_cell, m, c, ln, ci, ch, args, keys) + for m, c, ln, ci, ch in tasks + ] for done, (fut, t) in enumerate(zip(futures, tasks), 1): try: records.extend(fut.result()) except Exception as err: # noqa: BLE001 -- partial results still get written; rerun resumes from cache failed += 1 - print(f" FAIL {t[0]} {t[1]} len={t[2]} chunk={t[3]}: {type(err).__name__}: {err}", flush=True) + print( + f" FAIL {t[0]} {t[1]} len={t[2]} chunk={t[3]}: {type(err).__name__}: {err}", + flush=True, + ) print(f" {done}/{len(tasks)} cells", flush=True) with (out_dir / "records.jsonl").open("w") as fh: @@ -382,10 +437,25 @@ def main() -> None: for model in models: for length in lengths: for cond in conditions: - sub = [r for r in records if r["model"] == model and r["length"] == length and r["cond"] == cond] + sub = [ + r + for r in records + if r["model"] == model + and r["length"] == length + and r["cond"] == cond + ] if sub: - cells.append({"model": model, "length": length, "condition": cond, **aggregate(sub, *MODELS[model])}) - (out_dir / "summary.json").write_text(json.dumps({"args": vars(args), "cells": cells}, indent=1)) + cells.append( + { + "model": model, + "length": length, + "condition": cond, + **aggregate(sub, *MODELS[model]), + } + ) + (out_dir / "summary.json").write_text( + json.dumps({"args": vars(args), "cells": cells}, indent=1) + ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) writer.writeheader() @@ -399,7 +469,13 @@ def main() -> None: row = f"{cond:<18}" for model in models: cell = next( - (c for c in cells if c["model"] == model and c["length"] == length and c["condition"] == cond), + ( + c + for c in cells + if c["model"] == model + and c["length"] == length + and c["condition"] == cond + ), None, ) if cell: diff --git a/packages/snapcompact/research/exp07_readtax.py b/packages/snapcompact/research/exp07_readtax.py index 1a69432fb..9a6b18e09 100644 --- a/packages/snapcompact/research/exp07_readtax.py +++ b/packages/snapcompact/research/exp07_readtax.py @@ -43,8 +43,13 @@ from run import CACHE, FONTS, QA_CACHE, RESULTS, load_prompt, sha8 # noqa: E402 MODELS = {"gpt-5.5": (2.0, 16.0), "google/gemini-3.5-flash": (0.6, 4.0)} LENGTHS = (50, 150) CONDITIONS = ( - "baseline", "effort-low", "effort-minimal", "no-transcribe", - "locate-then-answer", "locate-low", "locate-none", + "baseline", + "effort-low", + "effort-minimal", + "no-transcribe", + "locate-then-answer", + "locate-low", + "locate-none", ) FONT, VARIANT = "6x10", "sent" @@ -63,10 +68,19 @@ def ensure_png(chunk_text: str, size: int) -> Path: return png -def timed_call(keys: dict, model: str, messages: list[dict], max_tokens: int, effort: str | None) -> dict: +def timed_call( + keys: dict, model: str, messages: list[dict], max_tokens: int, effort: str | None +) -> dict: t0 = time.monotonic() - text, usage, stop = llm_complete(keys, model, messages, max_tokens=max_tokens, effort=effort) - return {"text": text, "usage": usage, "stop": stop, "latency_s": round(time.monotonic() - t0, 2)} + text, usage, stop = llm_complete( + keys, model, messages, max_tokens=max_tokens, effort=effort + ) + return { + "text": text, + "usage": usage, + "stop": stop, + "latency_s": round(time.monotonic() - t0, 2), + } def probe_min_effort(keys: dict) -> str | None: @@ -74,9 +88,16 @@ def probe_min_effort(keys: dict) -> str | None: for effort in ("minimal", "none"): try: llm_complete( - keys, "gpt-5.5", - [{"role": "user", "content": [{"text": "Reply with the single word OK."}]}], - max_tokens=64, effort=effort, + keys, + "gpt-5.5", + [ + { + "role": "user", + "content": [{"text": "Reply with the single word OK."}], + } + ], + max_tokens=64, + effort=effort, ) return effort except SystemExit as err: @@ -84,9 +105,19 @@ def probe_min_effort(keys: dict) -> str | None: return None -def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> list[dict]: - args, flow, paras, offsets, keys = ctx["args"], ctx["flow"], ctx["paras"], ctx["offsets"], ctx["keys"] - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) +def run_cell_chunk( + model: str, cond: str, start: int, end: int, ctx: dict +) -> list[dict]: + args, flow, paras, offsets, keys = ( + ctx["args"], + ctx["flow"], + ctx["paras"], + ctx["offsets"], + ctx["keys"], + ) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] chunk_text = flow[start:end] @@ -108,20 +139,33 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li ] if cond.startswith("locate"): - turn2_effort = {"locate-then-answer": None, "locate-low": "low", "locate-none": "none"}[cond] + turn2_effort = { + "locate-then-answer": None, + "locate-low": "low", + "locate-none": "none", + }[cond] locate_msgs = qa_messages("exp07-locate.md") locate = cached( - model, "exp07-locate", {"messages": locate_msgs, "effort": "low"}, + model, + "exp07-locate", + {"messages": locate_msgs, "effort": "low"}, lambda: timed_call(keys, model, locate_msgs, args.max_tokens, "low"), args.fresh, ) - usage_rows.append(("locate", {**locate["usage"], "latency_s": locate.get("latency_s", 0)})) + usage_rows.append( + ("locate", {**locate["usage"], "latency_s": locate.get("latency_s", 0)}) + ) answer_msgs = locate_msgs + [ {"role": "assistant", "content": [{"text": locate["text"]}]}, - {"role": "user", "content": [{"text": load_prompt("exp07-answer-bands.md")}]}, + { + "role": "user", + "content": [{"text": load_prompt("exp07-answer-bands.md")}], + }, ] qa = cached( - model, "exp07-qa", {"cond": cond, "messages": answer_msgs, "effort": turn2_effort}, + model, + "exp07-qa", + {"cond": cond, "messages": answer_msgs, "effort": turn2_effort}, lambda: timed_call(keys, model, answer_msgs, args.max_tokens, turn2_effort), args.fresh, ) @@ -130,7 +174,9 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li effort = cond.removeprefix("effort-") if cond.startswith("effort-") else None messages = qa_messages(prompt_file) qa = cached( - model, "exp07-qa", {"cond": cond, "messages": messages, "effort": effort}, + model, + "exp07-qa", + {"cond": cond, "messages": messages, "effort": effort}, lambda: timed_call(keys, model, messages, args.max_tokens, effort), args.fresh, ) @@ -164,11 +210,18 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out # Per-chunk latency = sum over phases (locate + qa for the two-turn protocol). - chunk_lat = [sum(u.get("latency_s", 0) for u in r["usage"]) for r in records if "usage" in r] + chunk_lat = [ + sum(u.get("latency_s", 0) for u in r["usage"]) for r in records if "usage" in r + ] return { "n": n, "em": sum(r["em"] for r in records) / n, @@ -212,10 +265,17 @@ def main() -> None: "openai": load_env_key("OPENAI_API_KEY", args.env), "openrouter": load_env_key("OPENROUTER_API_KEY", args.env), } - min_effort = probe_min_effort(keys) if "effort-minimal" in conditions and "gpt-5.5" in models else None + min_effort = ( + probe_min_effort(keys) + if "effort-minimal" in conditions and "gpt-5.5" in models + else None + ) print(f"lowest gpt-5.5 effort: {min_effort or 'unavailable -> condition skipped'}") if "effort-minimal" in conditions: - conditions = [f"effort-{min_effort}" if c == "effort-minimal" and min_effort else c for c in conditions] + conditions = [ + f"effort-{min_effort}" if c == "effort-minimal" and min_effort else c + for c in conditions + ] conditions = [c for c in conditions if c != "effort-minimal"] budget = capacity(FONTS[FONT], args.size)[2] @@ -224,20 +284,34 @@ def main() -> None: for length in lengths: paras = all_paras[:length] flow, offsets = squad.build_flow(paras) - ctx = {"args": args, "flow": flow, "paras": paras, "offsets": offsets, "keys": keys, "length": length} + ctx = { + "args": args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + } for model in models: for cond in conditions: - if cond in (f"effort-{min_effort}", "locate-none") and model != "gpt-5.5": + if ( + cond in (f"effort-{min_effort}", "locate-none") + and model != "gpt-5.5" + ): continue if cond == "locate-none" and min_effort != "none": continue for start in range(0, len(flow), budget): - tasks.append((model, cond, start, min(start + budget, len(flow)), ctx)) + tasks.append( + (model, cond, start, min(start + budget, len(flow)), ctx) + ) print(f"grid: {len(tasks)} chunk tasks") records: list[dict] = [] with ThreadPoolExecutor(args.workers) as pool: - futures = [pool.submit(run_cell_chunk, m, c, s, e, ctx) for m, c, s, e, ctx in tasks] + futures = [ + pool.submit(run_cell_chunk, m, c, s, e, ctx) for m, c, s, e, ctx in tasks + ] for i, fut in enumerate(futures, 1): records.extend(fut.result()) print(f" {i}/{len(tasks)} tasks", flush=True) @@ -250,11 +324,26 @@ def main() -> None: for model in models: for length in lengths: for cond in conditions: - sub = [r for r in records if r["model"] == model and r["length"] == length and r["cond"] == cond] + sub = [ + r + for r in records + if r["model"] == model + and r["length"] == length + and r["cond"] == cond + ] if not sub: continue - cells.append({"model": model, "length": length, "condition": cond, **aggregate(sub, *MODELS[model])}) - (out_dir / "summary.json").write_text(json.dumps({"args": vars(args), "cells": cells}, indent=1)) + cells.append( + { + "model": model, + "length": length, + "condition": cond, + **aggregate(sub, *MODELS[model]), + } + ) + (out_dir / "summary.json").write_text( + json.dumps({"args": vars(args), "cells": cells}, indent=1) + ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) writer.writeheader() diff --git a/packages/snapcompact/research/exp08_foveate.py b/packages/snapcompact/research/exp08_foveate.py index 235765e79..2b1cc394f 100644 --- a/packages/snapcompact/research/exp08_foveate.py +++ b/packages/snapcompact/research/exp08_foveate.py @@ -122,7 +122,9 @@ def locate_phrase(chunk: str, phrase: str) -> tuple[int, int] | None: return None -def merge_bands(bands: list[tuple[int, int]], max_row: int, pad: int) -> list[tuple[int, int]]: +def merge_bands( + bands: list[tuple[int, int]], max_row: int, pad: int +) -> list[tuple[int, int]]: """Pad by `pad`, clamp to [1, max_row], merge overlapping/adjacent bands.""" padded = sorted((max(1, a - pad), min(max_row, b + pad)) for a, b in bands) merged: list[tuple[int, int]] = [] @@ -134,10 +136,14 @@ def merge_bands(bands: list[tuple[int, int]], max_row: int, pad: int) -> list[tu return merged -def zoom_renders(chunk_text: str, bands: list[tuple[int, int]], arch_cols: int) -> list[tuple[tuple[int, int], Path]]: +def zoom_renders( + chunk_text: str, bands: list[tuple[int, int]], arch_cols: int +) -> list[tuple[tuple[int, int], Path]]: """Slice each band's rows from the chunk and render at ZOOM_FONT; oversized bands split.""" zcfg = FONTS[ZOOM_FONT] - max_rows = capacity(zcfg, ZOOM_SIZES[-1])[2] // arch_cols # archive rows per zoom page + max_rows = ( + capacity(zcfg, ZOOM_SIZES[-1])[2] // arch_cols + ) # archive rows per zoom page out = [] for a, b in bands: pieces = [(s, min(s + max_rows - 1, b)) for s in range(a, b + 1, max_rows)] @@ -145,7 +151,10 @@ def zoom_renders(chunk_text: str, bands: list[tuple[int, int]], arch_cols: int) txt = chunk_text[(pa - 1) * arch_cols : pb * arch_cols] if not txt.strip(): continue - size = next((s for s in ZOOM_SIZES if capacity(zcfg, s)[2] >= len(txt)), ZOOM_SIZES[-1]) + size = next( + (s for s in ZOOM_SIZES if capacity(zcfg, s)[2] >= len(txt)), + ZOOM_SIZES[-1], + ) png = CACHE / f"exp08-zoom-{ZOOM_FONT}-{size}-{sha8(txt)}.png" if not png.exists() or png.stat().st_size == 0: atomic_png(render(txt, zcfg, CACHE, size, "bw"), png) @@ -153,10 +162,20 @@ def zoom_renders(chunk_text: str, bands: list[tuple[int, int]], arch_cols: int) return out -def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> list[dict]: +def run_cell_chunk( + model: str, cond: str, start: int, end: int, ctx: dict +) -> list[dict]: """One (model, condition, chunk): archive QA turn, optional zoom turn, merge, score.""" - args, flow, paras, offsets, keys = ctx["args"], ctx["flow"], ctx["paras"], ctx["offsets"], ctx["keys"] - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + args, flow, paras, offsets, keys = ( + ctx["args"], + ctx["flow"], + ctx["paras"], + ctx["offsets"], + ctx["keys"], + ) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] chunk_text = flow[start:end] @@ -165,7 +184,10 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li cfg = FONTS[ARCHIVE_FONT] cols, rows, _ = capacity(cfg, args.size) - png = CACHE / f"exp08-arch-{ARCHIVE_FONT}-{variant}-{sha8(chunk_text, str(args.size))}.png" + png = ( + CACHE + / f"exp08-arch-{ARCHIVE_FONT}-{variant}-{sha8(chunk_text, str(args.size))}.png" + ) if not png.exists() or png.stat().st_size == 0: atomic_png(render(chunk_text, cfg, CACHE, args.size, variant), png) @@ -181,9 +203,21 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li } ] qa1 = cached( - model, "exp08-qa1", {"messages": messages, "effort": args.effort}, - lambda: dict(zip(("text", "usage", "stop"), - llm_complete(keys, model, messages, max_tokens=args.max_tokens, effort=args.effort))), + model, + "exp08-qa1", + {"messages": messages, "effort": args.effort}, + lambda: dict( + zip( + ("text", "usage", "stop"), + llm_complete( + keys, + model, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), + ) + ), args.fresh, ) usage_rows = [("qa1", qa1["usage"])] @@ -196,11 +230,15 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li m = _ZOOM_PHRASE.search(a) text = m.group(1) if m else None if text is None and re.search(r"(?i)\bzoom\b", a) and not parse_zoom(a): - text = re.sub(r"(?i)^.*?\bzoom\b[:\s]*", "", a).strip("\"'\u201c\u201d ") + text = re.sub(r"(?i)^.*?\bzoom\b[:\s]*", "", a).strip( + "\"'\u201c\u201d " + ) if text and len(text.split()) >= 2: anchor = text span = locate_phrase(chunk_text, anchor) - zoom_req.append((span[0] // cols + 1, span[1] // cols + 1) if span else None) + zoom_req.append( + (span[0] // cols + 1, span[1] // cols + 1) if span else None + ) elif re.search(r"(?i)\bzoom\b", a): zoom_req.append(parse_zoom(a)) # rows fallback else: @@ -222,15 +260,29 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li for (a, b), zpng in zooms: z_content.append({"text": f"Zoom of archive rows {a}-{b}:"}) z_content.append({"image_path": zpng}) - z_content.append({"text": "\n".join(f"{i + 1}. {questions[i]['q']}" for i in pending)}) + z_content.append( + {"text": "\n".join(f"{i + 1}. {questions[i]['q']}" for i in pending)} + ) messages2 = messages + [ {"role": "assistant", "content": [{"text": qa1["text"]}]}, {"role": "user", "content": z_content}, ] qa2 = cached( - model, "exp08-qa2", {"messages": messages2, "effort": args.effort}, - lambda: dict(zip(("text", "usage", "stop"), - llm_complete(keys, model, messages2, max_tokens=args.max_tokens, effort=args.effort))), + model, + "exp08-qa2", + {"messages": messages2, "effort": args.effort}, + lambda: dict( + zip( + ("text", "usage", "stop"), + llm_complete( + keys, + model, + messages2, + max_tokens=args.max_tokens, + effort=args.effort, + ), + ) + ), args.fresh, ) usage_rows.append(("qa2", qa2["usage"])) @@ -265,7 +317,9 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li def _phase_cost(us: list[dict], price_in: float, price_out: float) -> float: tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r")} - return (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok["out"] / 1e6 * price_out + return ( + tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"] + ) / 1e6 * price_in + tok["out"] / 1e6 * price_out def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: @@ -274,11 +328,20 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out zoomed = sum(r["zoomed"] for r in records) - zoom_chunks = sum(1 for r in records if "usage" in r and any(u["phase"] == "qa2" for u in r["usage"])) + zoom_chunks = sum( + 1 + for r in records + if "usage" in r and any(u["phase"] == "qa2" for u in r["usage"]) + ) return { "n": n, "em": sum(r["em"] for r in records) / n, @@ -293,7 +356,9 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: "cost_in_usd": round(cost_in, 4), "cost_out_usd": round(cost_out, 4), "cost_usd": round(cost_in + cost_out, 4), - "cost_zoom_usd": round(_phase_cost([u for u in us if u["phase"] == "qa2"], price_in, price_out), 4), + "cost_zoom_usd": round( + _phase_cost([u for u in us if u["phase"] == "qa2"], price_in, price_out), 4 + ), } @@ -332,17 +397,30 @@ def main() -> None: for length in lengths: paras = all_paras[:length] flow, offsets = squad.build_flow(paras) - ctx = {"args": args, "flow": flow, "paras": paras, "offsets": offsets, "keys": keys, "length": length} + ctx = { + "args": args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + } for model in models: for cond in conditions: for start in range(0, len(flow), budget): - tasks.append((model, cond, start, min(start + budget, len(flow)), ctx)) - print(f"grid: {len(models)} models x {len(lengths)} lengths x {len(conditions)} conditions = {len(tasks)} chunk tasks") + tasks.append( + (model, cond, start, min(start + budget, len(flow)), ctx) + ) + print( + f"grid: {len(models)} models x {len(lengths)} lengths x {len(conditions)} conditions = {len(tasks)} chunk tasks" + ) records: list[dict] = [] done = 0 with ThreadPoolExecutor(args.workers) as pool: - futures = [pool.submit(run_cell_chunk, m, c, s, e, ctx) for m, c, s, e, ctx in tasks] + futures = [ + pool.submit(run_cell_chunk, m, c, s, e, ctx) for m, c, s, e, ctx in tasks + ] for fut in futures: records.extend(fut.result()) done += 1 @@ -356,10 +434,21 @@ def main() -> None: for model in models: for length in lengths: for cond in conditions: - sub = [r for r in records if r["model"] == model and r["length"] == length and r["cond"] == cond] + sub = [ + r + for r in records + if r["model"] == model + and r["length"] == length + and r["cond"] == cond + ] if not sub: continue - cell = {"model": model, "length": length, "condition": cond, **aggregate(sub, *MODELS[model])} + cell = { + "model": model, + "length": length, + "condition": cond, + **aggregate(sub, *MODELS[model]), + } base = BASELINE.get((model, length)) if base: cell["base_f1"] = base[0] @@ -367,7 +456,9 @@ def main() -> None: cell["d_f1"] = round(cell["f1"] - base[0], 4) cell["d_cost_usd"] = round(cell["cost_usd"] - base[2], 4) cells.append(cell) - (out_dir / "summary.json").write_text(json.dumps({"args": vars(args), "cells": cells}, indent=1)) + (out_dir / "summary.json").write_text( + json.dumps({"args": vars(args), "cells": cells}, indent=1) + ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) writer.writeheader() @@ -384,5 +475,3 @@ def main() -> None: if __name__ == "__main__": main() - - diff --git a/packages/snapcompact/research/exp09_cacheappend.py b/packages/snapcompact/research/exp09_cacheappend.py index 1d67070e0..8f6e68614 100644 --- a/packages/snapcompact/research/exp09_cacheappend.py +++ b/packages/snapcompact/research/exp09_cacheappend.py @@ -53,18 +53,35 @@ MAX_STEPS = 4 PROBE_TAIL = "Reply with exactly the word OK and nothing else." -def call(keys: dict, model: str, messages: list[dict], system: str | None = None, max_tokens: int = 32768) -> dict: +def call( + keys: dict, + model: str, + messages: list[dict], + system: str | None = None, + max_tokens: int = 32768, +) -> dict: t0 = time.monotonic() - text, usage, stop = llm_complete(keys, model, messages, system=system, max_tokens=max_tokens) - return {"text": text, "usage": usage, "stop": stop, "secs": round(time.monotonic() - t0, 2)} + text, usage, stop = llm_complete( + keys, model, messages, system=system, max_tokens=max_tokens + ) + return { + "text": text, + "usage": usage, + "stop": stop, + "secs": round(time.monotonic() - t0, 2), + } def usd(u: dict, p_in: float, p_out: float) -> float: - return (u.get("in", 0) + 0.1 * u.get("cache_r", 0)) / 1e6 * p_in + u.get("out", 0) / 1e6 * p_out + return (u.get("in", 0) + 0.1 * u.get("cache_r", 0)) / 1e6 * p_in + u.get( + "out", 0 + ) / 1e6 * p_out def usd_nocache(u: dict, p_in: float, p_out: float) -> float: - return (u.get("in", 0) + u.get("cache_r", 0)) / 1e6 * p_in + u.get("out", 0) / 1e6 * p_out + return (u.get("in", 0) + u.get("cache_r", 0)) / 1e6 * p_in + u.get( + "out", 0 + ) / 1e6 * p_out def render_pages(flow: str, size: int) -> tuple[list[tuple[int, int, Path]], dict]: @@ -93,21 +110,39 @@ def render_pages(flow: str, size: int) -> tuple[list[tuple[int, int, Path]], dic def prefix_messages(k: int, pages: list, cols: int, rows: int) -> list[dict]: """Append-only prefix: frame + pages 1..k, each ACKed. Byte-stable across steps.""" msgs = [ - {"role": "user", "content": [{"text": load_prompt("exp09-frame.md").format(cols=cols, rows=rows)}, {"image_path": pages[0][2]}]}, + { + "role": "user", + "content": [ + {"text": load_prompt("exp09-frame.md").format(cols=cols, rows=rows)}, + {"image_path": pages[0][2]}, + ], + }, {"role": "assistant", "content": [{"text": ACK}]}, ] for i in range(1, k): - msgs.append({"role": "user", "content": [{"text": load_prompt("exp09-page.md").format(page=i + 1)}, {"image_path": pages[i][2]}]}) + msgs.append( + { + "role": "user", + "content": [ + {"text": load_prompt("exp09-page.md").format(page=i + 1)}, + {"image_path": pages[i][2]}, + ], + } + ) msgs.append({"role": "assistant", "content": [{"text": ACK}]}) return msgs -def step_questions(paras: list, offsets: list, end: int, qpc: int, seed: int) -> tuple[list[dict], str]: +def step_questions( + paras: list, offsets: list, end: int, qpc: int, seed: int +) -> tuple[list[dict], str]: qs = squad.sample_chunk_questions(paras, offsets, 0, end, qpc, seed) return qs, "\n".join(f"{i + 1}. {q['q']}" for i, q in enumerate(qs)) -def score_records(model: str, regime: str, step: int, questions: list[dict], text: str) -> list[dict]: +def score_records( + model: str, regime: str, step: int, questions: list[dict], text: str +) -> list[dict]: answers = squad.parse_numbered(text, len(questions)) return [ { @@ -137,7 +172,14 @@ def common_prefix_len(a: str, b: str) -> int: def run_model(model: str, ctx: dict) -> dict: - args, keys, flow, paras, offsets, pages = ctx["args"], ctx["keys"], ctx["flow"], ctx["paras"], ctx["offsets"], ctx["pages"] + args, keys, flow, paras, offsets, pages = ( + ctx["args"], + ctx["keys"], + ctx["flow"], + ctx["paras"], + ctx["offsets"], + ctx["pages"], + ) p_in, p_out = MODELS[model] cols, rows, _ = capacity(FONTS[FONT], args.size) K = len(pages) @@ -150,11 +192,19 @@ def run_model(model: str, ctx: dict) -> dict: end = pages[k - 1][1] questions, q_block = step_questions(paras, offsets, end, args.qpc, args.seed) msgs = prefix_messages(k, pages, cols, rows) + [ - {"role": "user", "content": [{"text": load_prompt("exp09-qa.md").format(questions=q_block)}]} + { + "role": "user", + "content": [ + {"text": load_prompt("exp09-qa.md").format(questions=q_block)} + ], + } ] qa = cached( - model, "exp09-A-qa", {"step": k, "messages": msgs}, - lambda: call(keys, model, msgs, max_tokens=args.max_tokens), args.fresh, + model, + "exp09-A-qa", + {"step": k, "messages": msgs}, + lambda: call(keys, model, msgs, max_tokens=args.max_tokens), + args.fresh, ) recs = score_records(model, "append-optical", k, questions, qa["text"]) recs[0]["usage"] = [{"phase": "qa", **qa["usage"]}] @@ -164,27 +214,48 @@ def run_model(model: str, ctx: dict) -> dict: u = qa["usage"] steps.append( { - "model": model, "regime": "append-optical", "step": k, "n": len(recs), + "model": model, + "regime": "append-optical", + "step": k, + "n": len(recs), "f1": round(sum(r["f1"] for r in recs) / len(recs), 3), - "write_in": 0, "write_out": 0, "write_secs": 0.0, "write_cost": 0.0, - "qa_in": u["in"], "qa_cache_r": u["cache_r"], "qa_out": u["out"], - "qa_reasoning": u.get("reasoning", 0), "qa_secs": qa["secs"], - "step_cost": round(cost, 4), "step_cost_nocache": round(usd_nocache(u, p_in, p_out), 4), + "write_in": 0, + "write_out": 0, + "write_secs": 0.0, + "write_cost": 0.0, + "qa_in": u["in"], + "qa_cache_r": u["cache_r"], + "qa_out": u["out"], + "qa_reasoning": u.get("reasoning", 0), + "qa_secs": qa["secs"], + "step_cost": round(cost, 4), + "step_cost_nocache": round(usd_nocache(u, p_in, p_out), 4), "cum_cost": round(cum_a, 4), } ) - print(f" {model} A step {k}: in={u['in']} cache_r={u['cache_r']} out={u['out']} f1={steps[-1]['f1']}", flush=True) + print( + f" {model} A step {k}: in={u['in']} cache_r={u['cache_r']} out={u['out']} f1={steps[-1]['f1']}", + flush=True, + ) # --- Cache probe: identical multi-image prefix twice in a row (disk cache bypassed via call index) --- - probe_msgs = prefix_messages(K, pages, cols, rows) + [{"role": "user", "content": [{"text": PROBE_TAIL}]}] + probe_msgs = prefix_messages(K, pages, cols, rows) + [ + {"role": "user", "content": [{"text": PROBE_TAIL}]} + ] probe = [] for i in (1, 2, 3): r = cached( - model, "exp09-probe", {"call": i, "messages": probe_msgs}, - lambda: call(keys, model, probe_msgs, max_tokens=args.max_tokens), args.fresh, + model, + "exp09-probe", + {"call": i, "messages": probe_msgs}, + lambda: call(keys, model, probe_msgs, max_tokens=args.max_tokens), + args.fresh, ) probe.append({"call": i, **r["usage"], "secs": r["secs"]}) - print(f" {model} probe call {i}: in={r['usage']['in']} cache_r={r['usage']['cache_r']}", flush=True) + print( + f" {model} probe call {i}: in={r['usage']['in']} cache_r={r['usage']['cache_r']}", + flush=True, + ) # --- Regime B: rewrite-compact (fresh summary each step = write path) --- cum_b = 0.0 @@ -194,24 +265,46 @@ def run_model(model: str, ctx: dict) -> dict: text_k = flow[:end] questions, q_block = step_questions(paras, offsets, end, args.qpc, args.seed) sm = cached( - model, "exp09-B-sum", {"step": k, "chunk": sha8(text_k)}, + model, + "exp09-B-sum", + {"step": k, "chunk": sha8(text_k)}, lambda: call( - keys, model, - session_frame(text_k) + [{"role": "user", "content": [{"text": agent_prompt("compaction-summary.md")}]}], - system=agent_prompt("summarization-system.md"), max_tokens=args.max_tokens, + keys, + model, + session_frame(text_k) + + [ + { + "role": "user", + "content": [{"text": agent_prompt("compaction-summary.md")}], + } + ], + system=agent_prompt("summarization-system.md"), + max_tokens=args.max_tokens, ), args.fresh, ) summaries.append(sm["text"]) qa_msgs = [ - {"role": "user", "content": [{"text": load_prompt("qa-text.md").format(context=sm["text"])}, {"text": q_block}]} + { + "role": "user", + "content": [ + {"text": load_prompt("qa-text.md").format(context=sm["text"])}, + {"text": q_block}, + ], + } ] qa = cached( - model, "exp09-B-qa", {"step": k, "summary": sm["text"], "q": q_block}, - lambda: call(keys, model, qa_msgs, max_tokens=args.max_tokens), args.fresh, + model, + "exp09-B-qa", + {"step": k, "summary": sm["text"], "q": q_block}, + lambda: call(keys, model, qa_msgs, max_tokens=args.max_tokens), + args.fresh, ) recs = score_records(model, "rewrite-compact", k, questions, qa["text"]) - recs[0]["usage"] = [{"phase": "summarize", **sm["usage"]}, {"phase": "qa", **qa["usage"]}] + recs[0]["usage"] = [ + {"phase": "summarize", **sm["usage"]}, + {"phase": "qa", **qa["usage"]}, + ] records += recs w_cost = usd(sm["usage"], p_in, p_out) q_cost = usd(qa["usage"], p_in, p_out) @@ -219,41 +312,75 @@ def run_model(model: str, ctx: dict) -> dict: su, qu = sm["usage"], qa["usage"] steps.append( { - "model": model, "regime": "rewrite-compact", "step": k, "n": len(recs), + "model": model, + "regime": "rewrite-compact", + "step": k, + "n": len(recs), "f1": round(sum(r["f1"] for r in recs) / len(recs), 3), - "write_in": su["in"] + su["cache_r"], "write_out": su["out"], - "write_secs": sm["secs"], "write_cost": round(w_cost, 4), - "qa_in": qu["in"], "qa_cache_r": qu["cache_r"], "qa_out": qu["out"], - "qa_reasoning": qu.get("reasoning", 0), "qa_secs": qa["secs"], + "write_in": su["in"] + su["cache_r"], + "write_out": su["out"], + "write_secs": sm["secs"], + "write_cost": round(w_cost, 4), + "qa_in": qu["in"], + "qa_cache_r": qu["cache_r"], + "qa_out": qu["out"], + "qa_reasoning": qu.get("reasoning", 0), + "qa_secs": qa["secs"], "step_cost": round(w_cost + q_cost, 4), - "step_cost_nocache": round(usd_nocache(su, p_in, p_out) + usd_nocache(qu, p_in, p_out), 4), + "step_cost_nocache": round( + usd_nocache(su, p_in, p_out) + usd_nocache(qu, p_in, p_out), 4 + ), "cum_cost": round(cum_b, 4), } ) - print(f" {model} B step {k}: write {su['in']}+{su['cache_r']}c->{su['out']} ({sm['secs']}s) f1={steps[-1]['f1']}", flush=True) + print( + f" {model} B step {k}: write {su['in']}+{su['cache_r']}c->{su['out']} ({sm['secs']}s) f1={steps[-1]['f1']}", + flush=True, + ) # Write-path determinism of regime B: re-run the step-1 summarize with identical payload (fresh key). det = cached( - model, "exp09-B-sum-det", {"step": 1, "chunk": sha8(flow[: pages[0][1]])}, + model, + "exp09-B-sum-det", + {"step": 1, "chunk": sha8(flow[: pages[0][1]])}, lambda: call( - keys, model, - session_frame(flow[: pages[0][1]]) + [{"role": "user", "content": [{"text": agent_prompt("compaction-summary.md")}]}], - system=agent_prompt("summarization-system.md"), max_tokens=args.max_tokens, + keys, + model, + session_frame(flow[: pages[0][1]]) + + [ + { + "role": "user", + "content": [{"text": agent_prompt("compaction-summary.md")}], + } + ], + system=agent_prompt("summarization-system.md"), + max_tokens=args.max_tokens, ), args.fresh, ) b_det = { "identical": det["text"] == summaries[0], "common_prefix_chars": common_prefix_len(det["text"], summaries[0]), - "len_a": len(summaries[0]), "len_b": len(det["text"]), + "len_a": len(summaries[0]), + "len_b": len(det["text"]), } # Cross-step summary prefix stability (the thing the prompt cache would need). step_stability = [ - {"steps": f"{k}->{k + 1}", "common_prefix_chars": common_prefix_len(summaries[k - 1], summaries[k]), - "len_prev": len(summaries[k - 1]), "len_next": len(summaries[k])} + { + "steps": f"{k}->{k + 1}", + "common_prefix_chars": common_prefix_len(summaries[k - 1], summaries[k]), + "len_prev": len(summaries[k - 1]), + "len_next": len(summaries[k]), + } for k in range(1, K) ] - return {"records": records, "steps": steps, "probe": probe, "b_determinism": b_det, "b_step_stability": step_stability} + return { + "records": records, + "steps": steps, + "probe": probe, + "b_determinism": b_det, + "b_step_stability": step_stability, + } def aggregate(records: list[dict], p_in: float, p_out: float) -> dict: @@ -262,7 +389,10 @@ def aggregate(records: list[dict], p_in: float, p_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } cost_in = (tok["in"] + 0.1 * tok["cache_r"]) / 1e6 * p_in cost_out = tok["out"] / 1e6 * p_out return { @@ -304,9 +434,18 @@ def main() -> None: flow, offsets = squad.build_flow(paras) pages, a_det = render_pages(flow, args.size) pages = pages[:MAX_STEPS] - print(f"flow {len(flow)} chars -> {len(pages)} pages (K={len(pages)} steps); render deterministic: {a_det['deterministic']}") + print( + f"flow {len(flow)} chars -> {len(pages)} pages (K={len(pages)} steps); render deterministic: {a_det['deterministic']}" + ) - ctx = {"args": args, "keys": keys, "flow": flow, "paras": paras, "offsets": offsets, "pages": pages} + ctx = { + "args": args, + "keys": keys, + "flow": flow, + "paras": paras, + "offsets": offsets, + "pages": pages, + } with ThreadPoolExecutor(min(2, len(models))) as pool: results = dict(zip(models, pool.map(lambda m: run_model(m, ctx), models))) @@ -327,7 +466,12 @@ def main() -> None: sub = [r for r in records if r["model"] == model and r["cond"] == cond] final_step = max(r["step"] for r in sub) fin = [r for r in sub if r["step"] == final_step] - cell = {"model": model, "length": LENGTH, "condition": cond, **aggregate(sub, *MODELS[model])} + cell = { + "model": model, + "length": LENGTH, + "condition": cond, + **aggregate(sub, *MODELS[model]), + } cell["final_step_f1"] = round(sum(r["f1"] for r in fin) / len(fin), 3) cell["final_step_em"] = round(sum(r["em"] for r in fin) / len(fin), 3) cells.append(cell) @@ -348,7 +492,9 @@ def main() -> None: } (out_dir / "summary.json").write_text(json.dumps(summary, indent=1, default=str)) - print("\n== per-step (qa_in / qa_cache_r / write_cost / step_cost / cum_cost / f1) ==") + print( + "\n== per-step (qa_in / qa_cache_r / write_cost / step_cost / cum_cost / f1) ==" + ) for s in steps: print( f"{s['model']:<26} {s['regime']:<16} k={s['step']} in={s['qa_in']:>6} cache_r={s['qa_cache_r']:>6} " @@ -357,7 +503,9 @@ def main() -> None: print("\n== cache probe (same multi-image prefix twice) ==") for m in models: for p in results[m]["probe"]: - print(f"{m:<26} call {p['call']}: in={p['in']:>6} cache_r={p['cache_r']:>6} secs={p['secs']}") + print( + f"{m:<26} call {p['call']}: in={p['in']:>6} cache_r={p['cache_r']:>6} secs={p['secs']}" + ) print(f"\ndataset -> {out_dir}/records.jsonl, steps.csv, matrix.csv, summary.json") diff --git a/packages/snapcompact/research/exp10_profiles.py b/packages/snapcompact/research/exp10_profiles.py index 127edcd5f..c43b6734b 100644 --- a/packages/snapcompact/research/exp10_profiles.py +++ b/packages/snapcompact/research/exp10_profiles.py @@ -61,10 +61,10 @@ BASELINE = { # img-6x10-sent from results/optimal-{gpt55,gemini}/matrix.csv # Phase A screening cells (length 150). Sibling winners + variant probe at # 8x13 + the 6x12 bridge. img-6x10-sent is the baseline -- not re-run. SCREEN = ( - "img-6x12-dim", # fable's winner - "img-8x13-bw", # opus's winner + "img-6x12-dim", # fable's winner + "img-8x13-bw", # opus's winner "img-8x13-sent-dim", # kimi's winner - "img-8x13-dark-sent", # glm's winner + "img-8x13-dark-sent", # glm's winner "img-8x13-sent", "img-8x13-dim", "img-6x12-sent", @@ -110,11 +110,21 @@ def render_png(chunk_text: str, font: str, variant: str, size: int) -> Path: return png -def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> list[dict]: +def run_cell_chunk( + model: str, cond: str, start: int, end: int, ctx: dict +) -> list[dict]: """One (model, condition, chunk): render carrier image, QA, score. Copied from final.run_cell_chunk, image conditions only.""" - args, flow, paras, offsets, keys = ctx["args"], ctx["flow"], ctx["paras"], ctx["offsets"], ctx["keys"] - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + args, flow, paras, offsets, keys = ( + ctx["args"], + ctx["flow"], + ctx["paras"], + ctx["offsets"], + ctx["keys"], + ) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] chunk_text = flow[start:end] @@ -133,11 +143,14 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li } ] qa = cached( - model, {"messages": messages, "extra": None, "effort": None}, + model, + {"messages": messages, "extra": None, "effort": None}, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, model, messages, max_tokens=args.max_tokens, effort=None), + llm_complete( + keys, model, messages, max_tokens=args.max_tokens, effort=None + ), ) ), args.fresh, @@ -170,8 +183,13 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out return { "n": n, @@ -200,8 +218,12 @@ class Runner: paras = self.all_paras[:length] flow, offsets = squad.build_flow(paras) self.ctxs[length] = { - "args": self.args, "flow": flow, "paras": paras, - "offsets": offsets, "keys": self.keys, "length": length, + "args": self.args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": self.keys, + "length": length, } return self.ctxs[length] @@ -224,11 +246,30 @@ class Runner: self.records.extend(fut.result()) def cell(self, model: str, length: int, cond: str) -> dict | None: - sub = [r for r in self.records if r["model"] == model and r["length"] == length and r["cond"] == cond] - return {"model": model, "length": length, "condition": cond, **aggregate(sub, *MODELS[model])} if sub else None + sub = [ + r + for r in self.records + if r["model"] == model and r["length"] == length and r["cond"] == cond + ] + return ( + { + "model": model, + "length": length, + "condition": cond, + **aggregate(sub, *MODELS[model]), + } + if sub + else None + ) def cells_for(self, model: str, length: int) -> list[dict]: - conds = sorted({r["cond"] for r in self.records if r["model"] == model and r["length"] == length}) + conds = sorted( + { + r["cond"] + for r in self.records + if r["model"] == model and r["length"] == length + } + ) return [c for cond in conds if (c := self.cell(model, length, cond))] @@ -258,7 +299,9 @@ def main() -> None: runner.run([(m, 150, c) for m in models for c in SCREEN], "A screen") # -- phase B: density ladder at each model's best variant --------------- - ladder = [(m, 150, "img-6x10-sent") for m in models] # baseline; free via shared cache + ladder = [ + (m, 150, "img-6x10-sent") for m in models + ] # baseline; free via shared cache best_variant = {} for m in models: top = max(runner.cells_for(m, 150), key=lambda c: c["f1"]) @@ -269,12 +312,20 @@ def main() -> None: # -- phase C: cross-profile transfer (each model on the other's optimum) - top150 = {m: max(runner.cells_for(m, 150), key=lambda c: c["f1"]) for m in models} - cross = [(other, 150, top150[m]["condition"]) for m in models for other in models if other != m] + cross = [ + (other, 150, top150[m]["condition"]) + for m in models + for other in models + if other != m + ] runner.run(cross, "C cross") top150 = {m: max(runner.cells_for(m, 150), key=lambda c: c["f1"]) for m in models} # -- phase D: confirm top cell at lengths 50 and 250 --------------------- - runner.run([(m, ln, top150[m]["condition"]) for m in models for ln in (50, 250)], "D confirm") + runner.run( + [(m, ln, top150[m]["condition"]) for m in models for ln in (50, 250)], + "D confirm", + ) # -- outputs -------------------------------------------------------------- with (OUT_DIR / "records.jsonl").open("w") as fh: @@ -316,16 +367,25 @@ def main() -> None: "f1_se_at_150": round(top["f1_se"], 4), "cost_usd_at_150": top["cost_usd"], "confirm": { - str(ln): {"f1": round(c["f1"], 4), "se": round(c["f1_se"], 4), "cost_usd": c["cost_usd"]} - for ln, c in confirm.items() if c + str(ln): { + "f1": round(c["f1"], 4), + "se": round(c["f1_se"], 4), + "cost_usd": c["cost_usd"], + } + for ln, c in confirm.items() + if c }, }, "baseline_img_6x10_sent_f1_at_150": BASELINE[(m, 150)][0], - "transfer_f1_on_other_model_at_150": round(transfer["f1"], 4) if transfer else None, + "transfer_f1_on_other_model_at_150": round(transfer["f1"], 4) + if transfer + else None, } (OUT_DIR / "profiles.json").write_text(json.dumps(profiles, indent=1)) (OUT_DIR / "summary.json").write_text( - json.dumps({"args": vars(args), "best_variant": best_variant, "cells": cells}, indent=1) + json.dumps( + {"args": vars(args), "best_variant": best_variant, "cells": cells}, indent=1 + ) ) # -- console report ------------------------------------------------------- @@ -333,14 +393,18 @@ def main() -> None: for m in models: print(f"\n== {m} (length 150 screening, sorted by F1) ==") base_f1, base_se, base_cost = BASELINE[(m, 150)] - print(f"{'condition':<22}{'n':>5}{'EM':>7}{'F1':>7}{'se':>7}{'abst':>6}{'cost$':>8}{'dF1':>8}") + print( + f"{'condition':<22}{'n':>5}{'EM':>7}{'F1':>7}{'se':>7}{'abst':>6}{'cost$':>8}{'dF1':>8}" + ) for c in sorted(runner.cells_for(m, 150), key=lambda c: -c["f1"]): spend += c["cost_usd"] print( f"{c['condition']:<22}{c['n']:>5}{c['em']:>7.3f}{c['f1']:>7.3f}{c['f1_se']:>7.3f}" f"{c['abstained']:>6}{c['cost_usd']:>8.3f}{c['f1'] - base_f1:>+8.3f}" ) - print(f"{'img-6x10-sent [base]':<22}{'':>5}{'':>7}{base_f1:>7.3f}{base_se:>7.3f}{'':>6}{base_cost:>8.3f}{0:>+8.3f}") + print( + f"{'img-6x10-sent [base]':<22}{'':>5}{'':>7}{base_f1:>7.3f}{base_se:>7.3f}{'':>6}{base_cost:>8.3f}{0:>+8.3f}" + ) for ln in (50, 250): c = runner.cell(m, ln, top150[m]["condition"]) if c: diff --git a/packages/snapcompact/research/exp11_memhier.py b/packages/snapcompact/research/exp11_memhier.py index a3badacc3..558aa1886 100644 --- a/packages/snapcompact/research/exp11_memhier.py +++ b/packages/snapcompact/research/exp11_memhier.py @@ -31,7 +31,16 @@ import squad # noqa: E402 from bdf import capacity, render # noqa: E402 from final import MODELS, aggregate, cached, session_frame # noqa: E402 from providers import llm_complete, load_env_key # noqa: E402 -from run import CACHE, FONTS, QA_CACHE, RESULTS, TEXT_CHUNK, agent_prompt, load_prompt, sha8 # noqa: E402 +from run import ( + CACHE, + FONTS, + QA_CACHE, + RESULTS, + TEXT_CHUNK, + agent_prompt, + load_prompt, + sha8, +) # noqa: E402 L2_FONT, L2_VAR = "6x10", "sent" APX_FONT, APX_VAR = "5x8", "sent" @@ -64,16 +73,28 @@ def render_pages(text: str, font: str, var: str, size: int) -> list[Path]: return pages -def gen_summary(model: str, keys: dict, l3_text: str, max_tokens: int, fresh: bool) -> dict: +def gen_summary( + model: str, keys: dict, l3_text: str, max_tokens: int, fresh: bool +) -> dict: return cached( - model, "exp11-summary", {"chunk": l3_text}, + model, + "exp11-summary", + {"chunk": l3_text}, lambda: dict( zip( ("text", "usage", "stop"), llm_complete( - keys, model, + keys, + model, session_frame(l3_text) - + [{"role": "user", "content": [{"text": agent_prompt("compaction-summary.md")}]}], + + [ + { + "role": "user", + "content": [ + {"text": agent_prompt("compaction-summary.md")} + ], + } + ], system=agent_prompt("summarization-system.md"), max_tokens=max_tokens, ), @@ -83,7 +104,14 @@ def gen_summary(model: str, keys: dict, l3_text: str, max_tokens: int, fresh: bo ) -def context_blocks(cond: str, summary: str, l2_pages: list[Path], apx_pages: list[Path], l1_text: str, size: int) -> list[dict]: +def context_blocks( + cond: str, + summary: str, + l2_pages: list[Path], + apx_pages: list[Path], + l1_text: str, + size: int, +) -> list[dict]: cols, rows, _ = capacity(FONTS[L2_FONT], size) apx_note = "" if cond == "hier-appendix": @@ -92,22 +120,37 @@ def context_blocks(cond: str, summary: str, l2_pages: list[Path], apx_pages: lis frame = load_prompt("exp11-qa-hier.md").format( appendix_note=apx_note, n_pages=len(l2_pages), cols=cols, rows=rows ) - blocks: list[dict] = [{"text": frame}, {"text": f"TIER 3 — SUMMARY OF OLDEST THIRD:\n\n{summary}"}] + blocks: list[dict] = [ + {"text": frame}, + {"text": f"TIER 3 — SUMMARY OF OLDEST THIRD:\n\n{summary}"}, + ] if cond == "hier-appendix": for i, p in enumerate(apx_pages): - blocks.append({"text": f"TIER 3 appendix image {i + 1}/{len(apx_pages)} (same oldest text as dense bitmap):"}) + blocks.append( + { + "text": f"TIER 3 appendix image {i + 1}/{len(apx_pages)} (same oldest text as dense bitmap):" + } + ) blocks.append({"image_path": p}) for i, p in enumerate(l2_pages): - blocks.append({"text": f"TIER 2 page {i + 1}/{len(l2_pages)} (middle third as bitmap):"}) + blocks.append( + {"text": f"TIER 2 page {i + 1}/{len(l2_pages)} (middle third as bitmap):"} + ) blocks.append({"image_path": p}) - blocks.append({"text": f"TIER 1 — VERBATIM NEWEST THIRD:\n\n\n{l1_text}\n"}) + blocks.append( + { + "text": f"TIER 1 — VERBATIM NEWEST THIRD:\n\n\n{l1_text}\n" + } + ) return blocks def run_chunk(model: str, cond: str, start: int, end: int, cell: dict) -> list[dict]: """One QA call: shared hierarchical context + this chunk's question batch.""" args, keys, flow = cell["args"], cell["keys"], cell["flow"] - questions = squad.sample_chunk_questions(cell["paras"], cell["offsets"], start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + cell["paras"], cell["offsets"], start, end, args.qpc, args.seed + ) if not questions: return [] q_block = "\n".join(f"{i + 1}. {q['q']}" for i, q in enumerate(questions)) @@ -118,11 +161,24 @@ def run_chunk(model: str, cond: str, start: int, end: int, cell: dict) -> list[d } ] qa = cached( - model, "exp11-qa", {"cond": cond, "length": cell["length"], "messages": messages, "effort": args.effort}, + model, + "exp11-qa", + { + "cond": cond, + "length": cell["length"], + "messages": messages, + "effort": args.effort, + }, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, model, messages, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + model, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -205,7 +261,11 @@ def main() -> None: b1, b2 = tier_bounds(offsets, len(flow)) l3, l2, l1 = flow[:b1], flow[b1:b2], flow[b2:] l2_pages = render_pages(l2, L2_FONT, L2_VAR, args.size) - apx_pages = render_pages(l3, APX_FONT, APX_VAR, args.size) if "hier-appendix" in conditions else [] + apx_pages = ( + render_pages(l3, APX_FONT, APX_VAR, args.size) + if "hier-appendix" in conditions + else [] + ) print( f"length {length}: flow={len(flow)} chars, tiers L3={len(l3)} L2={len(l2)} L1={len(l1)}, " f"l2_pages={len(l2_pages)} apx_pages={len(apx_pages)}" @@ -213,7 +273,9 @@ def main() -> None: for model in models: summ = gen_summary(model, keys, l3, args.max_tokens, args.fresh) if summ.get("stop") == "max_tokens": - raise SystemExit(f"summary truncated for {model} length {length}; raise --max-tokens") + raise SystemExit( + f"summary truncated for {model} length {length}; raise --max-tokens" + ) summary_usage[(model, length)] = summ["usage"] print(f" summary[{model}]: {len(summ['text'])} chars") cells[(model, length)] = { @@ -225,7 +287,9 @@ def main() -> None: "length": length, "bounds": (b1, b2), "blocks": { - cond: context_blocks(cond, summ["text"], l2_pages, apx_pages, l1, args.size) + cond: context_blocks( + cond, summ["text"], l2_pages, apx_pages, l1, args.size + ) for cond in conditions }, } @@ -234,7 +298,15 @@ def main() -> None: for (model, length), cell in cells.items(): for cond in conditions: for start in range(0, len(cell["flow"]), TEXT_CHUNK): - tasks.append((model, cond, start, min(start + TEXT_CHUNK, len(cell["flow"])), cell)) + tasks.append( + ( + model, + cond, + start, + min(start + TEXT_CHUNK, len(cell["flow"])), + cell, + ) + ) print(f"grid: {len(tasks)} QA tasks") records: list[dict] = [] @@ -252,7 +324,9 @@ def main() -> None: for r in records: key = (r["model"], r["length"], r["cond"]) if key not in charged and "usage" in r: - r["usage"].append({"phase": "summarize", **summary_usage[(r["model"], r["length"])]}) + r["usage"].append( + {"phase": "summarize", **summary_usage[(r["model"], r["length"])]} + ) charged.add(key) with (out_dir / "records.jsonl").open("w") as fh: @@ -263,16 +337,41 @@ def main() -> None: for model in models: for length in lengths: for cond in conditions: - sub = [r for r in records if r["model"] == model and r["length"] == length and r["cond"] == cond] + sub = [ + r + for r in records + if r["model"] == model + and r["length"] == length + and r["cond"] == cond + ] if not sub: continue - cell_rows.append({"model": model, "length": length, "condition": cond, **aggregate(sub, *MODELS[model])}) + cell_rows.append( + { + "model": model, + "length": length, + "condition": cond, + **aggregate(sub, *MODELS[model]), + } + ) for tier in ("L3", "L2", "L1"): tsub = [r for r in sub if r["tier"] == tier] if tsub: - tier_rows.append({"model": model, "length": length, "condition": cond, "tier": tier, **tier_stats(tsub)}) + tier_rows.append( + { + "model": model, + "length": length, + "condition": cond, + "tier": tier, + **tier_stats(tsub), + } + ) - (out_dir / "summary.json").write_text(json.dumps({"args": vars(args), "cells": cell_rows, "tiers": tier_rows}, indent=1)) + (out_dir / "summary.json").write_text( + json.dumps( + {"args": vars(args), "cells": cell_rows, "tiers": tier_rows}, indent=1 + ) + ) with (out_dir / "matrix.csv").open("w", newline="") as fh: w = csv.DictWriter(fh, fieldnames=list(cell_rows[0].keys())) w.writeheader() @@ -294,7 +393,9 @@ def main() -> None: f"{t['model']:<24} {t['length']:>4} {t['condition']:<14} {t['tier']} n={t['n']:<3} " f"EM={t['em']:.3f} F1={t['f1']:.3f} ±{t['f1_se']:.3f} abst={t['abstained']}" ) - print(f"\nresults -> {out_dir}/records.jsonl, matrix.csv, terciles.csv, summary.json") + print( + f"\nresults -> {out_dir}/records.jsonl, matrix.csv, terciles.csv, summary.json" + ) if __name__ == "__main__": diff --git a/packages/snapcompact/research/exp12_arbitrage.py b/packages/snapcompact/research/exp12_arbitrage.py index 51f382789..f52c2eff3 100644 --- a/packages/snapcompact/research/exp12_arbitrage.py +++ b/packages/snapcompact/research/exp12_arbitrage.py @@ -39,8 +39,15 @@ INSTR = "Reply with exactly: OK" # ---------------------------------------------------------------- part A: mine -MINE_DIRS = ("optimal-combined", "optimal-gpt55", "optimal-gemini", "optimal-fable", - "optimal-opus", "optimal-kimi", "optimal-glm") +MINE_DIRS = ( + "optimal-combined", + "optimal-gpt55", + "optimal-gemini", + "optimal-fable", + "optimal-opus", + "optimal-kimi", + "optimal-glm", +) def cond_budget(cond: str) -> int | None: @@ -86,18 +93,27 @@ def mine() -> tuple[list[dict], dict]: continue tok = qa["in"] + qa["cache_r"] + qa["cache_w"] chars = min(r["chunk"] + budget, len(flows[r["length"]])) - r["chunk"] - cell = agg.setdefault((r["model"], r["cond"]), {"chars": 0, "tok_in": 0, "chunks": 0}) + cell = agg.setdefault( + (r["model"], r["cond"]), {"chars": 0, "tok_in": 0, "chunks": 0} + ) cell["chars"] += chars cell["tok_in"] += tok cell["chunks"] += 1 - detail.setdefault(r["model"], {}).setdefault(r["cond"], []).append((chars, tok)) + detail.setdefault(r["model"], {}).setdefault(r["cond"], []).append( + (chars, tok) + ) rows = [] for (model, cond), c in sorted(agg.items()): - rows.append({ - "model": model, "cond": cond, "chunks": c["chunks"], "chars": c["chars"], - "tok_in_total": c["tok_in"], - "chars_per_tok": round(c["chars"] / c["tok_in"], 3), - }) + rows.append( + { + "model": model, + "cond": cond, + "chunks": c["chunks"], + "chars": c["chars"], + "tok_in_total": c["tok_in"], + "chars_per_tok": round(c["chars"] / c["tok_in"], 3), + } + ) return rows, detail @@ -111,11 +127,17 @@ RL_PREFIXES = ("x-ratelimit", "ratelimit", "retry-after") def post_h(url: str, body: dict, headers: dict, retries: int = 4) -> tuple[dict, dict]: payload = json.dumps(body).encode() - req = urllib.request.Request(url, data=payload, headers={"content-type": "application/json", **headers}) + req = urllib.request.Request( + url, data=payload, headers={"content-type": "application/json", **headers} + ) for attempt in range(retries + 1): try: with urllib.request.urlopen(req, timeout=600) as resp: - rl = {k.lower(): v for k, v in resp.headers.items() if k.lower().startswith(RL_PREFIXES)} + rl = { + k.lower(): v + for k, v in resp.headers.items() + if k.lower().startswith(RL_PREFIXES) + } return json.load(resp), rl except urllib.error.HTTPError as err: detail = err.read().decode(errors="replace")[:300] @@ -140,29 +162,57 @@ def probe_call(model: str, keys: dict, blocks: list[dict]) -> tuple[dict, dict]: if "text" in b: content.append({"type": "input_text", "text": b["text"]}) else: - content.append({"type": "input_image", - "image_url": f"data:image/png;base64,{png_b64(b['image_path'])}", - "detail": "original"}) - body = {"model": model, "input": [{"role": "user", "content": content}], - "max_output_tokens": 512, "store": False} - out, rl = post_h(OPENAI_URL, body, {"authorization": f"Bearer {keys['openai']}"}) + content.append( + { + "type": "input_image", + "image_url": f"data:image/png;base64,{png_b64(b['image_path'])}", + "detail": "original", + } + ) + body = { + "model": model, + "input": [{"role": "user", "content": content}], + "max_output_tokens": 512, + "store": False, + } + out, rl = post_h( + OPENAI_URL, body, {"authorization": f"Bearer {keys['openai']}"} + ) u = out.get("usage", {}) cached = (u.get("input_tokens_details") or {}).get("cached_tokens", 0) - usage = {"in": u.get("input_tokens", 0), "cached": cached, "out": u.get("output_tokens", 0)} + usage = { + "in": u.get("input_tokens", 0), + "cached": cached, + "out": u.get("output_tokens", 0), + } return usage, rl content = [] for b in blocks: if "text" in b: content.append({"type": "text", "text": b["text"]}) else: - content.append({"type": "image_url", - "image_url": {"url": f"data:image/png;base64,{png_b64(b['image_path'])}"}}) - body = {"model": model, "messages": [{"role": "user", "content": content}], "max_tokens": 512} - out, rl = post_h(OPENROUTER_URL, body, {"authorization": f"Bearer {keys['openrouter']}"}) + content.append( + { + "type": "image_url", + "image_url": { + "url": f"data:image/png;base64,{png_b64(b['image_path'])}" + }, + } + ) + body = { + "model": model, + "messages": [{"role": "user", "content": content}], + "max_tokens": 512, + } + out, rl = post_h( + OPENROUTER_URL, body, {"authorization": f"Bearer {keys['openrouter']}"} + ) u = out.get("usage", {}) - usage = {"in": u.get("prompt_tokens", 0), - "cached": (u.get("prompt_tokens_details") or {}).get("cached_tokens", 0), - "out": u.get("completion_tokens", 0)} + usage = { + "in": u.get("prompt_tokens", 0), + "cached": (u.get("prompt_tokens_details") or {}).get("cached_tokens", 0), + "out": u.get("completion_tokens", 0), + } return usage, rl @@ -190,17 +240,23 @@ def run_probes(keys: dict, flow: str) -> dict: page_text = flow[:TEXT_CHUNK] out: dict = {"page_chars": len(page_text), "models": {}} for model in PROBE_MODELS: - steps = [("overhead-1", [{"text": INSTR}]), - ("text-page", [{"text": INSTR}, {"text": page_text}]), - ("overhead-2", [{"text": INSTR}])] - steps += [(f"img-{s}", [{"text": INSTR}, {"image_path": pngs[s]}]) for s in SIZES] + steps = [ + ("overhead-1", [{"text": INSTR}]), + ("text-page", [{"text": INSTR}, {"text": page_text}]), + ("overhead-2", [{"text": INSTR}]), + ] + steps += [ + (f"img-{s}", [{"text": INSTR}, {"image_path": pngs[s]}]) for s in SIZES + ] rows = [] for name, blocks in steps: usage, rl = probe_call(model, keys, blocks) row = {"step": name, "usage": usage, "ratelimit": rl, "t": time.time()} rows.append(row) - print(f" {model:>24} {name:<11} in={usage['in']:>6} (cached={usage['cached']}) " - f"out={usage['out']:>5} rl-remaining-tokens={rl.get('x-ratelimit-remaining-tokens', '-')}") + print( + f" {model:>24} {name:<11} in={usage['in']:>6} (cached={usage['cached']}) " + f"out={usage['out']:>5} rl-remaining-tokens={rl.get('x-ratelimit-remaining-tokens', '-')}" + ) out["models"][model] = rows return out @@ -253,9 +309,12 @@ def derive(mined: list[dict], detail: dict, probes: dict) -> dict: img = {} for s in SIZES: itok = by[f"img-{s}"]["in"] - overhead - img[s] = {"image_tokens": itok, "page_chars": caps[s], - "chars_per_img_tok": round(caps[s] / itok, 3), - "tok_per_megapixel": round(itok / (s * s / 1e6), 1)} + img[s] = { + "image_tokens": itok, + "page_chars": caps[s], + "chars_per_img_tok": round(caps[s] / itok, 3), + "tok_per_megapixel": round(itok / (s * s / 1e6), 1), + } cpt_text = page_chars / text_tok cpt_img = img[1568]["chars_per_img_tok"] stretch = cpt_img / cpt_text @@ -266,17 +325,31 @@ def derive(mined: list[dict], detail: dict, probes: dict) -> dict: "chars_per_text_tok": round(cpt_text, 3), "images": img, "window_stretch_6x10_1568": round(stretch, 3), - "chars_in_200k_window": {"text": int(200_000 * cpt_text), "img_6x10_1568": int(200_000 * cpt_img)}, + "chars_in_200k_window": { + "text": int(200_000 * cpt_text), + "img_6x10_1568": int(200_000 * cpt_img), + }, "breakeven_img_token_multiple": round(stretch, 3), - "input_cost_per_mchar": {"text": round(p_in / cpt_text, 4), "img_6x10_1568": round(p_in / cpt_img, 4)}, + "input_cost_per_mchar": { + "text": round(p_in / cpt_text, 4), + "img_6x10_1568": round(p_in / cpt_img, 4), + }, } - return {"mined": mined, "probes": probes, "derived": per_model, - "carrier_estimates": estimate_carriers(detail, per_model)} + return { + "mined": mined, + "probes": probes, + "derived": per_model, + "carrier_estimates": estimate_carriers(detail, per_model), + } def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--fresh", action="store_true", help="re-run API probes even if probes.json exists") + ap.add_argument( + "--fresh", + action="store_true", + help="re-run API probes even if probes.json exists", + ) ap.add_argument("--env", default="~/.env") args = ap.parse_args() @@ -286,16 +359,20 @@ def main() -> None: mined, detail = mine() print(f"mined {len(mined)} (model, cond) cells from {', '.join(MINE_DIRS)}") for r in mined: - print(f" {r['model']:>24} {r['cond']:<18} chunks={r['chunks']:>2} chars={r['chars']:>7} " - f"tok={r['tok_in_total']:>7} chars/tok={r['chars_per_tok']:>7.3f}") + print( + f" {r['model']:>24} {r['cond']:<18} chunks={r['chunks']:>2} chars={r['chars']:>7} " + f"tok={r['tok_in_total']:>7} chars/tok={r['chars_per_tok']:>7.3f}" + ) probes_path = OUT / "probes.json" if probes_path.exists() and not args.fresh: probes = json.loads(probes_path.read_text()) print("reusing probes.json (pass --fresh to re-run)") else: - keys = {"openai": load_env_key("OPENAI_API_KEY", args.env), - "openrouter": load_env_key("OPENROUTER_API_KEY", args.env)} + keys = { + "openai": load_env_key("OPENAI_API_KEY", args.env), + "openrouter": load_env_key("OPENROUTER_API_KEY", args.env), + } # 150 paragraphs -> flow ~90k chars, so every probe page (incl. 1568px / 40716 # chars) is completely full; image token cost is content-independent anyway # (verified: identical tok/megapixel at three different fill ratios). @@ -311,9 +388,11 @@ def main() -> None: tmp.replace(OUT / "measurements.json") print(f"\nwrote {OUT}/measurements.json") for model, d in measurements["derived"].items(): - print(f"{model}: text {d['chars_per_text_tok']} c/t | img-1568 " - f"{d['images'][1568 if 1568 in d['images'] else '1568']['chars_per_img_tok']} c/t | " - f"stretch {d['window_stretch_6x10_1568']}x | breakeven {d['breakeven_img_token_multiple']}x") + print( + f"{model}: text {d['chars_per_text_tok']} c/t | img-1568 " + f"{d['images'][1568 if 1568 in d['images'] else '1568']['chars_per_img_tok']} c/t | " + f"stretch {d['window_stretch_6x10_1568']}x | breakeven {d['breakeven_img_token_multiple']}x" + ) if __name__ == "__main__": diff --git a/packages/snapcompact/research/exp13_extractive.py b/packages/snapcompact/research/exp13_extractive.py index d23db7221..3c1960527 100644 --- a/packages/snapcompact/research/exp13_extractive.py +++ b/packages/snapcompact/research/exp13_extractive.py @@ -64,7 +64,12 @@ def cached(model: str, tag: str, payload: object, fn, fresh: bool) -> dict: def session_frame(chunk_text: str) -> list[dict]: return [ - {"role": "user", "content": [{"text": load_prompt("session-frame.md").format(context=chunk_text)}]}, + { + "role": "user", + "content": [ + {"text": load_prompt("session-frame.md").format(context=chunk_text)} + ], + }, {"role": "assistant", "content": [{"text": ACK}]}, ] @@ -75,8 +80,16 @@ def gold_survives(golds: list[str], extraction_norm: str) -> bool: def run_cell_chunk(model: str, start: int, end: int, ctx: dict) -> list[dict]: """One (model, chunk) unit: extract verbatim spans, QA over the extraction, score.""" - args, flow, paras, offsets, keys = ctx["args"], ctx["flow"], ctx["paras"], ctx["offsets"], ctx["keys"] - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + args, flow, paras, offsets, keys = ( + ctx["args"], + ctx["flow"], + ctx["paras"], + ctx["offsets"], + ctx["keys"], + ) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] chunk_text = flow[start:end] @@ -85,13 +98,17 @@ def run_cell_chunk(model: str, start: int, end: int, ctx: dict) -> list[dict]: extract_prompt = load_prompt("exp13-extract.md").format(budget=args.budget) gen = cached( - model, "exp13-extract", {"chunk": chunk_text, "budget": args.budget, "effort": args.extract_effort}, + model, + "exp13-extract", + {"chunk": chunk_text, "budget": args.budget, "effort": args.extract_effort}, lambda: dict( zip( ("text", "usage", "stop"), llm_complete( - keys, model, - session_frame(chunk_text) + [{"role": "user", "content": [{"text": extract_prompt}]}], + keys, + model, + session_frame(chunk_text) + + [{"role": "user", "content": [{"text": extract_prompt}]}], max_tokens=args.extract_max_tokens, effort=args.extract_effort, ), @@ -106,11 +123,16 @@ def run_cell_chunk(model: str, start: int, end: int, ctx: dict) -> list[dict]: messages = [ { "role": "user", - "content": [{"text": load_prompt("qa-text.md").format(context=extraction)}, {"text": q_block}], + "content": [ + {"text": load_prompt("qa-text.md").format(context=extraction)}, + {"text": q_block}, + ], } ] qa = cached( - model, "exp13-qa", {"messages": messages}, + model, + "exp13-qa", + {"messages": messages}, lambda: dict( zip( ("text", "usage", "stop"), @@ -150,8 +172,13 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out return { "n": n, @@ -174,11 +201,19 @@ def main() -> None: ap.add_argument("--lengths", default="50,150,250") ap.add_argument("--qpc", type=int, default=30) ap.add_argument("--seed", type=int, default=42) - ap.add_argument("--budget", type=int, default=8000, help="max extraction chars per chunk") - ap.add_argument("--max-tokens", type=int, default=32768, help="QA max tokens") - ap.add_argument("--extract-max-tokens", type=int, default=16384, help="extraction max tokens (budget+slack)") ap.add_argument( - "--extract-effort", default="low", + "--budget", type=int, default=8000, help="max extraction chars per chunk" + ) + ap.add_argument("--max-tokens", type=int, default=32768, help="QA max tokens") + ap.add_argument( + "--extract-max-tokens", + type=int, + default=16384, + help="extraction max tokens (budget+slack)", + ) + ap.add_argument( + "--extract-effort", + default="low", help="reasoning effort for the extraction call only; verbatim copying needs no deliberation " "(default-effort gemini burns ~16k reasoning tokens verifying quotes and truncates)", ) @@ -209,11 +244,20 @@ def main() -> None: for length in lengths: paras = all_paras[:length] flow, offsets = squad.build_flow(paras) - ctx = {"args": args, "flow": flow, "paras": paras, "offsets": offsets, "keys": keys, "length": length} + ctx = { + "args": args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + } for model in models: for start in range(0, len(flow), TEXT_CHUNK): tasks.append((model, start, min(start + TEXT_CHUNK, len(flow)), ctx)) - print(f"grid: {len(models)} models x {len(lengths)} lengths x 1 condition = {len(tasks)} chunk tasks") + print( + f"grid: {len(models)} models x {len(lengths)} lengths x 1 condition = {len(tasks)} chunk tasks" + ) records: list[dict] = [] done = 0 @@ -234,8 +278,17 @@ def main() -> None: sub = [r for r in records if r["model"] == model and r["length"] == length] if not sub: continue - cells.append({"model": model, "length": length, "condition": COND, **aggregate(sub, *MODELS[model])}) - (out_dir / "summary.json").write_text(json.dumps({"args": vars(args), "cells": cells}, indent=1)) + cells.append( + { + "model": model, + "length": length, + "condition": COND, + **aggregate(sub, *MODELS[model]), + } + ) + (out_dir / "summary.json").write_text( + json.dumps({"args": vars(args), "cells": cells}, indent=1) + ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) writer.writeheader() @@ -251,5 +304,3 @@ def main() -> None: if __name__ == "__main__": main() - - diff --git a/packages/snapcompact/research/exp14_bestgpt.py b/packages/snapcompact/research/exp14_bestgpt.py index 9ed35b88d..4448a5253 100644 --- a/packages/snapcompact/research/exp14_bestgpt.py +++ b/packages/snapcompact/research/exp14_bestgpt.py @@ -47,7 +47,9 @@ EXP = "exp14" OUT_DIR = RESULTS / f"{EXP}-bestgpt" MODEL = "gpt-5.5" PRICE_IN, PRICE_OUT = 2.0, 16.0 -FONT = FontCfg("8on16", "8x13", 8, 16) # exp01 winner: 8x13 glyphs, 16px patch-aligned pitch +FONT = FontCfg( + "8on16", "8x13", 8, 16 +) # exp01 winner: 8x13 glyphs, 16px patch-aligned pitch GUTTER = 3 # char cells between doc columns (as exp04) SCREEN_CELLS = "img-doc-8on16-bw@150,img-doc-8on16-sent@150,img-8on16-bw@150" _WHITE = (255, 255, 255) @@ -200,16 +202,30 @@ def atomic_save(img: Image.Image, png: Path) -> None: tmp.replace(png) -def qa_call(messages: list[dict], questions: list[dict], length: int, cond: str, - start: int, ctx: dict) -> list[dict]: +def qa_call( + messages: list[dict], + questions: list[dict], + length: int, + cond: str, + start: int, + ctx: dict, +) -> list[dict]: """One QA call + scoring; shared by doc and grid paths.""" args, keys = ctx["args"], ctx["keys"] qa = cached( - MODEL, f"{EXP}-qa", {"messages": messages, "effort": args.effort}, + MODEL, + f"{EXP}-qa", + {"messages": messages, "effort": args.effort}, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, MODEL, messages, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + MODEL, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -237,17 +253,23 @@ def qa_call(messages: list[dict], questions: list[dict], length: int, cond: str, return records -def run_doc_page(cond: str, length: int, page: tuple[int, int], ctx: dict) -> list[dict]: +def run_doc_page( + cond: str, length: int, page: tuple[int, int], ctx: dict +) -> list[dict]: args, paras, offsets = ctx["args"], ctx["paras"], ctx["offsets"] i, j = page start = offsets[i] end = offsets[j - 1] + len(paras[j - 1]["ctx"]) - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] variant = cond.removeprefix("img-doc-8on16-") lines = ctx["lines"][page] - page_key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size)) + page_key = sha8( + cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size) + ) png = CACHE / f"{EXP}-doc-{variant}-{page_key}.png" if not png.exists() or png.stat().st_size == 0: atomic_save(render_doc(lines, args.size, variant, CACHE), png) @@ -258,7 +280,11 @@ def run_doc_page(cond: str, length: int, page: tuple[int, int], ctx: dict) -> li { "role": "user", "content": [ - {"text": load_prompt("exp04-qa-image.md").format(col_w=col_w, rows=rows)}, + { + "text": load_prompt("exp04-qa-image.md").format( + col_w=col_w, rows=rows + ) + }, {"image_path": png}, {"text": q_block}, ], @@ -267,9 +293,13 @@ def run_doc_page(cond: str, length: int, page: tuple[int, int], ctx: dict) -> li return qa_call(messages, questions, length, cond, start, ctx) -def run_grid_chunk(cond: str, length: int, start: int, end: int, ctx: dict) -> list[dict]: +def run_grid_chunk( + cond: str, length: int, start: int, end: int, ctx: dict +) -> list[dict]: args, flow, paras, offsets = ctx["args"], ctx["flow"], ctx["paras"], ctx["offsets"] - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] chunk_text = flow[start:end] @@ -298,8 +328,13 @@ def aggregate(records: list[dict]) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * PRICE_IN + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * PRICE_IN + ) cost_out = tok["out"] / 1e6 * PRICE_OUT return { "n": n, @@ -318,18 +353,35 @@ def cell_label(cond: str, effort: str | None) -> str: return f"{cond}+eff-{effort}" if effort else cond -def write_outputs(records: list[dict], capacity_stats: dict, args_dict: dict) -> list[dict]: +def write_outputs( + records: list[dict], capacity_stats: dict, args_dict: dict +) -> list[dict]: with (OUT_DIR / "records.jsonl").open("w") as fh: for r in records: fh.write(json.dumps(r) + "\n") - cell_keys = sorted({(r["length"], r["cond"], r.get("effort")) for r in records}, - key=lambda k: (k[0], k[1], k[2] or "")) + cell_keys = sorted( + {(r["length"], r["cond"], r.get("effort")) for r in records}, + key=lambda k: (k[0], k[1], k[2] or ""), + ) cells = [] for length, cond, effort in cell_keys: - sub = [r for r in records if r["length"] == length and r["cond"] == cond and r.get("effort") == effort] - cells.append({"model": MODEL, "length": length, "condition": cell_label(cond, effort), **aggregate(sub)}) + sub = [ + r + for r in records + if r["length"] == length and r["cond"] == cond and r.get("effort") == effort + ] + cells.append( + { + "model": MODEL, + "length": length, + "condition": cell_label(cond, effort), + **aggregate(sub), + } + ) (OUT_DIR / "summary.json").write_text( - json.dumps({"args": args_dict, "capacity": capacity_stats, "cells": cells}, indent=1) + json.dumps( + {"args": args_dict, "capacity": capacity_stats, "cells": cells}, indent=1 + ) ) with (OUT_DIR / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) @@ -349,7 +401,9 @@ def main() -> None: ap.add_argument("--effort", default=None) ap.add_argument("--fresh", action="store_true") ap.add_argument("--render-only", action="store_true") - ap.add_argument("--report", action="store_true", help="re-aggregate existing records, no API") + ap.add_argument( + "--report", action="store_true", help="re-aggregate existing records, no API" + ) ap.add_argument("--env", default="~/.env") args = ap.parse_args() @@ -360,22 +414,28 @@ def main() -> None: cols, rows, grid_cap = capacity(FONT, args.size) col_w = (cols - GUTTER) // 2 max_lines = 2 * rows - print(f"8on16 @ {args.size}px: grid {cols}x{rows} = {grid_cap} chars; " - f"doc 2 x {col_w} cols + gutter {GUTTER}, {max_lines} line slots") + print( + f"8on16 @ {args.size}px: grid {cols}x{rows} = {grid_cap} chars; " + f"doc 2 x {col_w} cols + gutter {GUTTER}, {max_lines} line slots" + ) rec_path = OUT_DIR / "records.jsonl" existing: list[dict] = [] if rec_path.exists(): - existing = [json.loads(ln) for ln in rec_path.read_text().splitlines() if ln.strip()] + existing = [ + json.loads(ln) for ln in rec_path.read_text().splitlines() if ln.strip() + ] cap_path = OUT_DIR / "capacity.json" capacity_stats: dict = json.loads(cap_path.read_text()) if cap_path.exists() else {} if args.report: cells = write_outputs(existing, capacity_stats, vars(args)) for c in cells: - print(f"len {c['length']:<4} {c['condition']:<28} n={c['n']:<4} EM {c['em']:.3f} " - f"F1 {c['f1']:.3f} ±{c['f1_se']:.3f} ${c['cost_usd']:.3f} " - f"out={c['tok_out']} rsn={c['tok_reasoning']}") + print( + f"len {c['length']:<4} {c['condition']:<28} n={c['n']:<4} EM {c['em']:.3f} " + f"F1 {c['f1']:.3f} ±{c['f1_se']:.3f} ${c['cost_usd']:.3f} " + f"out={c['tok_out']} rsn={c['tok_reasoning']}" + ) return cell_specs = [] @@ -398,7 +458,9 @@ def main() -> None: flow, offsets = squad.build_flow(paras) pages = pack_pages(paras, col_w, max_lines) page_lines = {pg: layout_page(paras[pg[0] : pg[1]], col_w) for pg in pages} - page_chars = [offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] for i, j in pages] + page_chars = [ + offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] for i, j in pages + ] capacity_stats[str(length)] = { "doc_pages": len(pages), "mean_chars_page": round(sum(page_chars) / len(pages)), @@ -409,10 +471,19 @@ def main() -> None: "grid_pages": -(-len(flow) // grid_cap), } st = capacity_stats[str(length)] - print(f" len {length}: {st['doc_pages']} doc pages (mean {st['mean_chars_page']} chars, " - f"{round(100 * st['mean_chars_page'] / grid_cap)}% of grid {grid_cap}); " - f"grid {st['grid_pages']} pages; corpus {st['corpus_chars']}") - ctx = {"args": args, "paras": paras, "flow": flow, "offsets": offsets, "keys": keys, "lines": page_lines} + print( + f" len {length}: {st['doc_pages']} doc pages (mean {st['mean_chars_page']} chars, " + f"{round(100 * st['mean_chars_page'] / grid_cap)}% of grid {grid_cap}); " + f"grid {st['grid_pages']} pages; corpus {st['corpus_chars']}" + ) + ctx = { + "args": args, + "paras": paras, + "flow": flow, + "offsets": offsets, + "keys": keys, + "lines": page_lines, + } for cond, ln in cell_specs: if ln != length: continue @@ -421,7 +492,15 @@ def main() -> None: tasks.append(("doc", cond, length, pg, ctx)) else: for start in range(0, len(flow), grid_cap): - tasks.append(("grid", cond, length, (start, min(start + grid_cap, len(flow))), ctx)) + tasks.append( + ( + "grid", + cond, + length, + (start, min(start + grid_cap, len(flow))), + ctx, + ) + ) cap_path.write_text(json.dumps(capacity_stats, indent=1)) @@ -432,13 +511,22 @@ def main() -> None: if kind == "doc": variant = cond.removeprefix("img-doc-8on16-") i, j = unit - key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in ctx["paras"][i:j]]), str(args.size)) + key = sha8( + cond, + json.dumps([(p["title"], p["ctx"]) for p in ctx["paras"][i:j]]), + str(args.size), + ) png = CACHE / f"{EXP}-doc-{variant}-{key}.png" - atomic_save(render_doc(ctx["lines"][unit], args.size, variant, CACHE), png) + atomic_save( + render_doc(ctx["lines"][unit], args.size, variant, CACHE), png + ) else: variant = cond.removeprefix("img-8on16-") chunk_text = ctx["flow"][unit[0] : unit[1]] - png = CACHE / f"{EXP}-8on16-{variant}-{sha8(chunk_text, str(args.size))}.png" + png = ( + CACHE + / f"{EXP}-8on16-{variant}-{sha8(chunk_text, str(args.size))}.png" + ) atomic_save(render(chunk_text, FONT, CACHE, args.size, variant), png) print(f" sample: {png}") return @@ -452,7 +540,9 @@ def main() -> None: if kind == "doc": futures.append(pool.submit(run_doc_page, cond, length, unit, ctx)) else: - futures.append(pool.submit(run_grid_chunk, cond, length, unit[0], unit[1], ctx)) + futures.append( + pool.submit(run_grid_chunk, cond, length, unit[0], unit[1], ctx) + ) for fut in futures: new_records.extend(fut.result()) done += 1 @@ -460,14 +550,18 @@ def main() -> None: # merge: drop existing records for cells just re-run, keep everything else rerun = {(ln, cond, args.effort) for cond, ln in cell_specs} - kept = [r for r in existing if (r["length"], r["cond"], r.get("effort")) not in rerun] + kept = [ + r for r in existing if (r["length"], r["cond"], r.get("effort")) not in rerun + ] records = kept + new_records cells = write_outputs(records, capacity_stats, vars(args)) for c in cells: - print(f"len {c['length']:<4} {c['condition']:<28} n={c['n']:<4} EM {c['em']:.3f} " - f"F1 {c['f1']:.3f} ±{c['f1_se']:.3f} ${c['cost_usd']:.3f} " - f"out={c['tok_out']} rsn={c['tok_reasoning']}") + print( + f"len {c['length']:<4} {c['condition']:<28} n={c['n']:<4} EM {c['em']:.3f} " + f"F1 {c['f1']:.3f} ±{c['f1_se']:.3f} ${c['cost_usd']:.3f} " + f"out={c['tok_out']} rsn={c['tok_reasoning']}" + ) print(f"\n-> {OUT_DIR}/records.jsonl, matrix.csv, summary.json") diff --git a/packages/snapcompact/research/exp15_bestgemini.py b/packages/snapcompact/research/exp15_bestgemini.py index 1119b32e3..e45c3d98d 100644 --- a/packages/snapcompact/research/exp15_bestgemini.py +++ b/packages/snapcompact/research/exp15_bestgemini.py @@ -31,7 +31,16 @@ HERE = Path(__file__).resolve().parent sys.path.insert(0, str(HERE)) import squad # noqa: E402 -from bdf import _DARK, _DIMMED, _stopword_mask, FontCfg, capacity, parse_bdf, ensure_font, render # noqa: E402 +from bdf import ( + _DARK, + _DIMMED, + _stopword_mask, + FontCfg, + capacity, + parse_bdf, + ensure_font, + render, +) # noqa: E402 from providers import llm_complete, load_env_key # noqa: E402 from run import CACHE, QA_CACHE, RESULTS, load_prompt, sha8 # noqa: E402 @@ -143,7 +152,9 @@ def _doc_colors(lines: list[dict], variant: str) -> list[list[tuple[int, int, in n = len(ln["text"]) colors.append( [ - _DIMMED if dim is not None and dim[pos + k] else _DARK[sidx[pos + k] % 6] + _DIMMED + if dim is not None and dim[pos + k] + else _DARK[sidx[pos + k] % 6] for k in range(n) ] ) @@ -191,7 +202,14 @@ def render_doc(lines: list[dict], size: int, variant: str, cache: Path) -> Image # --- runners ----------------------------------------------------------------- -def _qa_call(model: str, messages: list[dict], questions: list[dict], ctx: dict, cond: str, start: int) -> list[dict]: +def _qa_call( + model: str, + messages: list[dict], + questions: list[dict], + ctx: dict, + cond: str, + start: int, +) -> list[dict]: args, keys = ctx["args"], ctx["keys"] qa = cached( model, @@ -199,7 +217,13 @@ def _qa_call(model: str, messages: list[dict], questions: list[dict], ctx: dict, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, model, messages, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + model, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -231,12 +255,16 @@ def run_doc_page(model: str, cond: str, page: tuple[int, int], ctx: dict) -> lis i, j = page start = offsets[i] end = offsets[j - 1] + len(paras[j - 1]["ctx"]) - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] variant = cond.removeprefix("img-doc-8on16-") lines = ctx["lines"][page] - page_key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size)) + page_key = sha8( + cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size) + ) png = CACHE / f"{EXP}-doc-8on16-{variant}-{page_key}.png" if not png.exists() or png.stat().st_size == 0: tmp = png.with_suffix(f".{os.getpid()}.tmp.png") @@ -249,7 +277,11 @@ def run_doc_page(model: str, cond: str, page: tuple[int, int], ctx: dict) -> lis { "role": "user", "content": [ - {"text": load_prompt("exp04-qa-image.md").format(col_w=col_w, rows=rows)}, + { + "text": load_prompt("exp04-qa-image.md").format( + col_w=col_w, rows=rows + ) + }, {"image_path": png}, {"text": q_block}, ], @@ -258,9 +290,13 @@ def run_doc_page(model: str, cond: str, page: tuple[int, int], ctx: dict) -> lis return _qa_call(model, messages, questions, ctx, cond, start) -def run_grid_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> list[dict]: +def run_grid_chunk( + model: str, cond: str, start: int, end: int, ctx: dict +) -> list[dict]: args, flow, paras, offsets = ctx["args"], ctx["flow"], ctx["paras"], ctx["offsets"] - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] chunk_text = flow[start:end] @@ -291,8 +327,13 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out return { "n": n, @@ -326,7 +367,9 @@ class Runner: col_w = (cols - GUTTER) // 2 pages = pack_pages(paras, col_w, 2 * rows) page_lines = {pg: layout_page(paras[pg[0] : pg[1]], col_w) for pg in pages} - page_chars = [offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] for i, j in pages] + page_chars = [ + offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] for i, j in pages + ] self.capacity_stats[length] = { "doc_pages": len(pages), "doc_mean_chars_page": round(sum(page_chars) / len(pages)), @@ -336,8 +379,14 @@ class Runner: "grid_pages": -(-len(flow) // grid_cap), } self.ctxs[length] = { - "args": self.args, "flow": flow, "paras": paras, "offsets": offsets, - "keys": self.keys, "length": length, "pages": pages, "lines": page_lines, + "args": self.args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": self.keys, + "length": length, + "pages": pages, + "lines": page_lines, } return self.ctxs[length] @@ -354,7 +403,12 @@ class Runner: grid_cap = capacity(FONT, self.args.size)[2] flow = ctx["flow"] for start in range(0, len(flow), grid_cap): - tasks.append((run_grid_chunk, (MODEL, cond, start, min(start + grid_cap, len(flow)), ctx))) + tasks.append( + ( + run_grid_chunk, + (MODEL, cond, start, min(start + grid_cap, len(flow)), ctx), + ) + ) if not tasks: return print(f"[{label}] {len(cells)} cells -> {len(tasks)} page/chunk tasks") @@ -366,7 +420,16 @@ class Runner: def cell(self, length: int, cond: str) -> dict | None: sub = [r for r in self.records if r["length"] == length and r["cond"] == cond] - return {"model": MODEL, "length": length, "condition": cond, **aggregate(sub, *PRICE)} if sub else None + return ( + { + "model": MODEL, + "length": length, + "condition": cond, + **aggregate(sub, *PRICE), + } + if sub + else None + ) def all_cells(self) -> list[dict]: keys = sorted({(r["length"], r["cond"]) for r in self.records}) @@ -400,8 +463,14 @@ def main() -> None: ap.add_argument("--effort", default=None) ap.add_argument("--fresh", action="store_true") ap.add_argument("--render-only", action="store_true") - ap.add_argument("--screen-only", action="store_true", help="skip the 50/250 confirmation phase") - ap.add_argument("--confirm-conds", default=None, help="comma list; default = screening F1 winner") + ap.add_argument( + "--screen-only", action="store_true", help="skip the 50/250 confirmation phase" + ) + ap.add_argument( + "--confirm-conds", + default=None, + help="comma list; default = screening F1 winner", + ) ap.add_argument("--env", default="~/.env") args = ap.parse_args() @@ -415,7 +484,9 @@ def main() -> None: cols, rows, grid_cap = capacity(FONT, args.size) col_w = (cols - GUTTER) // 2 - print(f"font 8on16: {cols} cols x {rows} rows; grid cap {grid_cap}; doc 2 x {col_w} + gutter {GUTTER}, {2 * rows} line slots") + print( + f"font 8on16: {cols} cols x {rows} rows; grid cap {grid_cap}; doc 2 x {col_w} + gutter {GUTTER}, {2 * rows} line slots" + ) runner = Runner(args, keys) @@ -425,21 +496,32 @@ def main() -> None: for cond in SCREEN_CONDS: if cond.startswith("img-doc-"): variant = cond.removeprefix("img-doc-8on16-") - key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in ctx["paras"][pg[0] : pg[1]]]), str(args.size)) + key = sha8( + cond, + json.dumps( + [(p["title"], p["ctx"]) for p in ctx["paras"][pg[0] : pg[1]]] + ), + str(args.size), + ) png = CACHE / f"{EXP}-doc-8on16-{variant}-{key}.png" img = render_doc(ctx["lines"][pg], args.size, variant, CACHE) else: variant = cond.removeprefix("img-8on16-") chunk = ctx["flow"][:grid_cap] - png = CACHE / f"{EXP}-grid-8on16-{variant}-{sha8(chunk, str(args.size))}.png" + png = ( + CACHE + / f"{EXP}-grid-8on16-{variant}-{sha8(chunk, str(args.size))}.png" + ) img = render(chunk, FONT, CACHE, args.size, variant) tmp = png.with_suffix(f".{os.getpid()}.tmp.png") img.save(tmp) tmp.replace(png) print(f" sample: {png}") for length, st in runner.capacity_stats.items(): - print(f" len {length}: {st['doc_pages']} doc pages, mean {st['doc_mean_chars_page']} chars/page " - f"(grid {st['grid_chars_page']} -> {st['grid_pages']} pages)") + print( + f" len {length}: {st['doc_pages']} doc pages, mean {st['doc_mean_chars_page']} chars/page " + f"(grid {st['grid_chars_page']} -> {st['grid_pages']} pages)" + ) return # Phase A: screen at 150 @@ -451,7 +533,11 @@ def main() -> None: # Phase B: confirm winner at 50 and 250 if not args.screen_only: - confirm = [w.strip() for w in args.confirm_conds.split(",")] if args.confirm_conds else [winner] + confirm = ( + [w.strip() for w in args.confirm_conds.split(",")] + if args.confirm_conds + else [winner] + ) runner.run([(ln, c) for ln in (50, 250) for c in confirm], "confirm@50/250") cells = runner.all_cells() @@ -460,13 +546,20 @@ def main() -> None: fh.write(json.dumps(r) + "\n") (OUT_DIR / "summary.json").write_text( json.dumps( - {"args": vars(args), "model": MODEL, "capacity": runner.capacity_stats, - "screen_winner": winner, "cells": cells}, + { + "args": vars(args), + "model": MODEL, + "capacity": runner.capacity_stats, + "screen_winner": winner, + "cells": cells, + }, indent=1, ) ) with (OUT_DIR / "matrix.csv").open("w", newline="") as fh: - writer = csv.DictWriter(fh, fieldnames=[k for k in cells[0].keys() if k != "doc_chars_page"]) + writer = csv.DictWriter( + fh, fieldnames=[k for k in cells[0].keys() if k != "doc_chars_page"] + ) writer.writeheader() writer.writerows(cells) diff --git a/packages/snapcompact/research/exp16_bestfable.py b/packages/snapcompact/research/exp16_bestfable.py index 8045f588c..57b49fbbd 100644 --- a/packages/snapcompact/research/exp16_bestfable.py +++ b/packages/snapcompact/research/exp16_bestfable.py @@ -39,7 +39,15 @@ HERE = Path(__file__).resolve().parent sys.path.insert(0, str(HERE)) import squad # noqa: E402 -from bdf import _DIMMED, _stopword_mask, FontCfg, capacity, ensure_font, parse_bdf, render # noqa: E402 +from bdf import ( + _DIMMED, + _stopword_mask, + FontCfg, + capacity, + ensure_font, + parse_bdf, + render, +) # noqa: E402 from providers import llm_complete, load_env_key # noqa: E402 from run import CACHE, QA_CACHE, RESULTS, load_prompt, sha8 # noqa: E402 @@ -47,7 +55,9 @@ MODEL = "claude-fable-5" PRICE = (10.0, 50.0) # $/M in, out FONTS = { "8on16": FontCfg("8on16", "8x13", 8, 16), # 8x13 glyphs, patch-aligned 16 px pitch - "6on7x14": FontCfg("6on7x14", "6x12", 7, 14), # 6x12 glyphs, 7x14 patch-aligned cell + "6on7x14": FontCfg( + "6on7x14", "6x12", 7, 14 + ), # 6x12 glyphs, 7x14 patch-aligned cell "6x12": FontCfg("6x12", "6x12", 6, 12), # fable's round-0 winner font } CONDITIONS = ("img-8on16-dim", "img-6on7x14-dim", "doc-6x12-dim", "doc-8on16-dim") @@ -57,8 +67,16 @@ _WHITE = (255, 255, 255) _BLACK = (0, 0, 0) # img-6x12-dim per length: (f1, se, cost); text ceiling: (f1, se, cost). -BASE_IMG = {50: (0.9556, 0.0348, 0.132), 150: (0.9113, 0.0244, 0.437), 250: (0.9233, 0.0163, 0.724)} -BASE_TEXT = {50: (0.9556, 0.0348, 0.144), 150: (0.9043, 0.0216, 0.498), 250: (0.9197, 0.0184, 0.734)} +BASE_IMG = { + 50: (0.9556, 0.0348, 0.132), + 150: (0.9113, 0.0244, 0.437), + 250: (0.9233, 0.0163, 0.724), +} +BASE_TEXT = { + 50: (0.9556, 0.0348, 0.144), + 150: (0.9043, 0.0216, 0.498), + 250: (0.9197, 0.0184, 0.734), +} def cached(model: str, tag: str, payload: object, fn, fresh: bool) -> dict: @@ -208,18 +226,28 @@ def render_doc(lines: list[dict], cfg: FontCfg, size: int, cache: Path) -> Image def qa_call(cond: str, messages: list[dict], ctx: dict) -> dict: args, keys = ctx["args"], ctx["keys"] return cached( - MODEL, f"exp16-qa-{cond}", {"messages": messages, "size": args.size, "effort": args.effort}, + MODEL, + f"exp16-qa-{cond}", + {"messages": messages, "size": args.size, "effort": args.effort}, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, MODEL, messages, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + MODEL, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, ) -def score(questions: list[dict], qa: dict, cond: str, start: int, ctx: dict) -> list[dict]: +def score( + questions: list[dict], qa: dict, cond: str, start: int, ctx: dict +) -> list[dict]: answers = squad.parse_numbered(qa["text"], len(questions)) records = [] for q, a in zip(questions, answers): @@ -245,7 +273,9 @@ def score(questions: list[dict], qa: dict, cond: str, start: int, ctx: dict) -> def run_grid_chunk(cond: str, start: int, end: int, ctx: dict) -> list[dict]: """Row-major grid cell: chunk the flow by capacity, one QA call per chunk.""" args, flow, paras, offsets = ctx["args"], ctx["flow"], ctx["paras"], ctx["offsets"] - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] chunk_text = flow[start:end] @@ -274,13 +304,17 @@ def run_doc_page(cond: str, page: tuple[int, int], ctx: dict) -> list[dict]: i, j = page start = offsets[i] end = offsets[j - 1] + len(paras[j - 1]["ctx"]) - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] _, font, _ = parse_condition(cond) cfg = FONTS[font] lines = ctx["lines"][cond][page] - page_key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size)) + page_key = sha8( + cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size) + ) png = CACHE / f"exp16-{cond}-{page_key}.png" if not png.exists() or png.stat().st_size == 0: atomic_save(render_doc(lines, cfg, args.size, CACHE), png) @@ -291,7 +325,11 @@ def run_doc_page(cond: str, page: tuple[int, int], ctx: dict) -> list[dict]: { "role": "user", "content": [ - {"text": load_prompt("exp04-qa-image.md").format(col_w=col_w, rows=rows)}, + { + "text": load_prompt("exp04-qa-image.md").format( + col_w=col_w, rows=rows + ) + }, {"image_path": png}, {"text": q_block}, ], @@ -306,8 +344,13 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out return { "n": n, @@ -333,7 +376,11 @@ def main() -> None: ap.add_argument("--max-tokens", type=int, default=32768) ap.add_argument("--effort", default=None) ap.add_argument("--fresh", action="store_true") - ap.add_argument("--render-only", action="store_true", help="render sample pages + capacity stats, no API") + ap.add_argument( + "--render-only", + action="store_true", + help="render sample pages + capacity stats, no API", + ) ap.add_argument("--env", default="~/.env") args = ap.parse_args() @@ -365,8 +412,13 @@ def main() -> None: col_w = (cols - GUTTER) // 2 pages = pack_pages(paras, col_w, 2 * rows) doc_pages[cond] = pages - page_lines[cond] = {pg: layout_page(paras[pg[0] : pg[1]], col_w) for pg in pages} - page_chars = [offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] for i, j in pages] + page_lines[cond] = { + pg: layout_page(paras[pg[0] : pg[1]], col_w) for pg in pages + } + page_chars = [ + offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] + for i, j in pages + ] capacity_stats[f"{cond}@{length}"] = { "pages": len(pages), "mean_chars_page": round(sum(page_chars) / len(pages)), @@ -381,8 +433,13 @@ def main() -> None: "corpus_chars": len(flow), } ctx = { - "args": args, "flow": flow, "paras": paras, "offsets": offsets, - "keys": keys, "length": length, "lines": page_lines, + "args": args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + "lines": page_lines, } for cond in conditions: kind, font, _ = parse_condition(cond) @@ -392,7 +449,9 @@ def main() -> None: else: budget = capacity(FONTS[font], args.size)[2] for start in range(0, len(flow), budget): - tasks.append(("img", cond, (start, min(start + budget, len(flow))), ctx)) + tasks.append( + ("img", cond, (start, min(start + budget, len(flow))), ctx) + ) for key, st in sorted(capacity_stats.items()): print( @@ -411,7 +470,11 @@ def main() -> None: col_w = (cols - GUTTER) // 2 i, j = pack_pages(paras, col_w, 2 * rows)[0] lines = layout_page(paras[i:j], col_w) - key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size)) + key = sha8( + cond, + json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), + str(args.size), + ) png = CACHE / f"exp16-{cond}-{key}.png" atomic_save(render_doc(lines, cfg, args.size, CACHE), png) else: @@ -419,7 +482,10 @@ def main() -> None: cap = capacity(cfg, args.size)[2] chunk_text = flow[:cap] _, _, variant = parse_condition(cond) - png = CACHE / f"exp16-{font}-{variant}-{sha8(chunk_text, str(args.size))}.png" + png = ( + CACHE + / f"exp16-{font}-{variant}-{sha8(chunk_text, str(args.size))}.png" + ) atomic_save(render(chunk_text, cfg, CACHE, args.size, variant), png) print(f" sample: {png}") return @@ -449,9 +515,18 @@ def main() -> None: sub = [r for r in records if r["length"] == length and r["cond"] == cond] if not sub: continue - cells.append({"model": MODEL, "length": length, "condition": cond, **aggregate(sub, *PRICE)}) + cells.append( + { + "model": MODEL, + "length": length, + "condition": cond, + **aggregate(sub, *PRICE), + } + ) (out_dir / "summary.json").write_text( - json.dumps({"args": vars(args), "capacity": capacity_stats, "cells": cells}, indent=1) + json.dumps( + {"args": vars(args), "capacity": capacity_stats, "cells": cells}, indent=1 + ) ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) @@ -461,7 +536,11 @@ def main() -> None: for c in cells: bi, bt = BASE_IMG.get(c["length"]), BASE_TEXT.get(c["length"]) comb_se = (c["f1_se"] ** 2 + bi[1] ** 2) ** 0.5 if bi else 0.0 - d_img = f"vs 6x12-dim {c['f1'] - bi[0]:+.3f} ({(c['f1'] - bi[0]) / comb_se:+.1f}se)" if bi else "" + d_img = ( + f"vs 6x12-dim {c['f1'] - bi[0]:+.3f} ({(c['f1'] - bi[0]) / comb_se:+.1f}se)" + if bi + else "" + ) d_txt = f" vs text {c['f1'] - bt[0]:+.3f}" if bt else "" flag = " ** beats text ceiling" if bt and c["f1"] > bt[0] else "" print( diff --git a/packages/snapcompact/research/exp17_bestopus.py b/packages/snapcompact/research/exp17_bestopus.py index 01ec418fd..f80b21b79 100644 --- a/packages/snapcompact/research/exp17_bestopus.py +++ b/packages/snapcompact/research/exp17_bestopus.py @@ -56,7 +56,11 @@ _BLACK = (0, 0, 0) _INK = (24, 24, 24) # near-black body text, like a printed page # claude-opus-4-8 baselines, results/optimal-combined/matrix.csv (qpc 30, seed 42): -BASELINE = {50: (0.9626, 0.0258, 0.143), 150: (0.8937, 0.0223, 0.380), 250: (0.8708, 0.0196, 0.559)} +BASELINE = { + 50: (0.9626, 0.0258, 0.143), + 150: (0.8937, 0.0223, 0.380), + 250: (0.8708, 0.0196, 0.559), +} TEXT_CEIL = {50: (0.9278, 0.195), 150: (0.9112, 0.637), 250: (0.9268, 0.938)} @@ -198,7 +202,11 @@ def render_unit_png(cond: str, unit: tuple[int, int], ctx: dict) -> Path: tmp.replace(png) else: i, j = unit - key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size)) + key = sha8( + cond, + json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), + str(args.size), + ) png = CACHE / f"exp17-{cond}-{key}.png" if not png.exists() or png.stat().st_size == 0: tmp = png.with_suffix(f".{uuid.uuid4().hex[:8]}.tmp.png") @@ -217,7 +225,9 @@ def run_unit(cond: str, unit: tuple[int, int], ctx: dict) -> list[dict]: i, j = unit start = offsets[i] end = offsets[j - 1] + len(paras[j - 1]["ctx"]) - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] q_block = "\n".join(f"{k + 1}. {q['q']}" for k, q in enumerate(questions)) @@ -226,7 +236,9 @@ def run_unit(cond: str, unit: tuple[int, int], ctx: dict) -> list[dict]: if kind == "grid": preamble = load_prompt("qa-image.md").format(cols=cols, rows=rows) else: - preamble = load_prompt("exp04-qa-image.md").format(col_w=(cols - GUTTER) // 2, rows=rows) + preamble = load_prompt("exp04-qa-image.md").format( + col_w=(cols - GUTTER) // 2, rows=rows + ) messages = [ { "role": "user", @@ -238,11 +250,19 @@ def run_unit(cond: str, unit: tuple[int, int], ctx: dict) -> list[dict]: } ] qa = cached( - MODEL, "exp17-qa", {"messages": messages, "effort": args.effort}, + MODEL, + "exp17-qa", + {"messages": messages, "effort": args.effort}, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, MODEL, messages, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + MODEL, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -275,8 +295,13 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out return { "n": n, @@ -302,8 +327,16 @@ def main() -> None: ap.add_argument("--max-tokens", type=int, default=32768) ap.add_argument("--effort", default=None) ap.add_argument("--fresh", action="store_true") - ap.add_argument("--render-only", action="store_true", help="capacity stats + first-page PNGs, no API") - ap.add_argument("--report", action="store_true", help="reprint matrix from accumulated records only") + ap.add_argument( + "--render-only", + action="store_true", + help="capacity stats + first-page PNGs, no API", + ) + ap.add_argument( + "--report", + action="store_true", + help="reprint matrix from accumulated records only", + ) ap.add_argument("--env", default="~/.env") args = ap.parse_args() @@ -320,29 +353,46 @@ def main() -> None: records: list[dict] = [] if not args.report: - keys = {} if args.render_only else {"anthropic": load_env_key("ANTHROPIC_API_KEY", args.env)} + keys = ( + {} + if args.render_only + else {"anthropic": load_env_key("ANTHROPIC_API_KEY", args.env)} + ) tasks = [] capacity_stats = {} for length in lengths: paras = all_paras[:length] flow, offsets = squad.build_flow(paras) ctx = { - "args": args, "flow": flow, "paras": paras, "offsets": offsets, - "keys": keys, "length": length, "lines": {}, + "args": args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + "lines": {}, } for cond in conditions: kind, cfg = parse_cond(cond) cols, rows, grid_cap = capacity(cfg, args.size) if kind == "grid": - units = [(s, min(s + grid_cap, len(flow))) for s in range(0, len(flow), grid_cap)] + units = [ + (s, min(s + grid_cap, len(flow))) + for s in range(0, len(flow), grid_cap) + ] chars = [e - s for s, e in units] else: col_w = (cols - GUTTER) // 2 pages = pack_pages(paras, col_w, 2 * rows) for pg in pages: - ctx["lines"][(cond, pg)] = layout_page(paras[pg[0] : pg[1]], col_w) + ctx["lines"][(cond, pg)] = layout_page( + paras[pg[0] : pg[1]], col_w + ) units = pages - chars = [offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] for i, j in pages] + chars = [ + offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] + for i, j in pages + ] capacity_stats[f"{cond}@{length}"] = { "pages": len(units), "mean_chars_page": round(sum(chars) / len(units)), @@ -381,7 +431,9 @@ def main() -> None: if records_path.exists(): with records_path.open() as fh: old = [json.loads(ln) for ln in fh if ln.strip()] - records = [r for r in old if (r["length"], r["cond"]) not in ran_cells] + records + records = [ + r for r in old if (r["length"], r["cond"]) not in ran_cells + ] + records with records_path.open("w") as fh: for r in records: fh.write(json.dumps(r) + "\n") @@ -393,16 +445,28 @@ def main() -> None: for length in sorted({r["length"] for r in records}): for cond in sorted({r["cond"] for r in records if r["length"] == length}): sub = [r for r in records if r["length"] == length and r["cond"] == cond] - cells.append({"model": MODEL, "length": length, "condition": cond, **aggregate(sub, *PRICES)}) + cells.append( + { + "model": MODEL, + "length": length, + "condition": cond, + **aggregate(sub, *PRICES), + } + ) (out_dir / "summary.json").write_text( - json.dumps({"args": vars(args), "baseline_img_8x13_bw": BASELINE, "cells": cells}, indent=1) + json.dumps( + {"args": vars(args), "baseline_img_8x13_bw": BASELINE, "cells": cells}, + indent=1, + ) ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) writer.writeheader() writer.writerows(cells) - print(f"\n{'len':<5}{'condition':<22}{'n':>5}{'EM':>7}{'F1':>7}{'+-se':>7}{'$':>8}{'d/se vs 8x13-bw':>17}") + print( + f"\n{'len':<5}{'condition':<22}{'n':>5}{'EM':>7}{'F1':>7}{'+-se':>7}{'$':>8}{'d/se vs 8x13-bw':>17}" + ) for c in cells: b_f1, b_se, b_cost = BASELINE[c["length"]] dse = (c["f1"] - b_f1) / ((c["f1_se"] ** 2 + b_se**2) ** 0.5 or 1) diff --git a/packages/snapcompact/research/exp18_bestkimi.py b/packages/snapcompact/research/exp18_bestkimi.py index 60c2f0601..61406f839 100644 --- a/packages/snapcompact/research/exp18_bestkimi.py +++ b/packages/snapcompact/research/exp18_bestkimi.py @@ -34,7 +34,16 @@ HERE = Path(__file__).resolve().parent sys.path.insert(0, str(HERE)) import squad # noqa: E402 -from bdf import _DARK, _DIMMED, FontCfg, _stopword_mask, capacity, ensure_font, parse_bdf, render # noqa: E402 +from bdf import ( + _DARK, + _DIMMED, + FontCfg, + _stopword_mask, + capacity, + ensure_font, + parse_bdf, + render, +) # noqa: E402 from providers import llm_complete, load_env_key # noqa: E402 from run import CACHE, QA_CACHE, RESULTS, load_prompt, sha8 # noqa: E402 @@ -199,7 +208,9 @@ def render_doc(lines: list[dict], cfg: FontCfg, size: int, cache: Path) -> Image # --- runner ----------------------------------------------------------------- -def doc_png(cond: str, paras: list[dict], lines: list[dict], cfg: FontCfg, size: int) -> Path: +def doc_png( + cond: str, paras: list[dict], lines: list[dict], cfg: FontCfg, size: int +) -> Path: key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in paras]), str(size)) png = CACHE / f"exp18-{cond}-{key}.png" if not png.exists() or png.stat().st_size == 0: @@ -213,7 +224,9 @@ def run_unit(cond: str, unit: dict, ctx: dict) -> list[dict]: """One (condition, page/chunk) unit: render carrier, QA, score.""" args, paras, offsets, keys = ctx["args"], ctx["paras"], ctx["offsets"], ctx["keys"] start, end = unit["start"], unit["end"] - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] kind, font, variant = CONDITIONS[cond] @@ -240,11 +253,19 @@ def run_unit(cond: str, unit: dict, ctx: dict) -> list[dict]: } ] qa = cached( - MODEL, "exp18-qa", {"messages": messages, "size": args.size, "effort": args.effort}, + MODEL, + "exp18-qa", + {"messages": messages, "size": args.size, "effort": args.effort}, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, MODEL, messages, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + MODEL, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -277,8 +298,13 @@ def aggregate(records: list[dict]) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * PRICE_IN + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * PRICE_IN + ) cost_out = tok["out"] / 1e6 * PRICE_OUT return { "n": n, @@ -304,8 +330,16 @@ def main() -> None: ap.add_argument("--max-tokens", type=int, default=32768) ap.add_argument("--effort", default=None) ap.add_argument("--fresh", action="store_true") - ap.add_argument("--render-only", action="store_true", help="render first page per cond + capacity stats, no API") - ap.add_argument("--report", action="store_true", help="re-aggregate (all units should hit cache)") + ap.add_argument( + "--render-only", + action="store_true", + help="render first page per cond + capacity stats, no API", + ) + ap.add_argument( + "--report", + action="store_true", + help="re-aggregate (all units should hit cache)", + ) ap.add_argument("--env", default="~/.env") args = ap.parse_args() @@ -330,7 +364,14 @@ def main() -> None: for length in lengths: paras = all_paras[:length] flow, offsets = squad.build_flow(paras) - ctx = {"args": args, "flow": flow, "paras": paras, "offsets": offsets, "keys": keys, "length": length} + ctx = { + "args": args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + } capacity_stats[length] = {"corpus_chars": len(flow), "conds": {}} for cond in conditions: kind, font, _ = CONDITIONS[cond] @@ -338,7 +379,8 @@ def main() -> None: cols, rows, grid_cap = capacity(cfg, args.size) if kind == "grid": units = [ - {"start": s, "end": min(s + grid_cap, len(flow))} for s in range(0, len(flow), grid_cap) + {"start": s, "end": min(s + grid_cap, len(flow))} + for s in range(0, len(flow), grid_cap) ] chars = [u["end"] - u["start"] for u in units] else: @@ -366,7 +408,9 @@ def main() -> None: for length, st in capacity_stats.items(): print(f"len {length}: corpus {st['corpus_chars']} chars") for cond, cs in st["conds"].items(): - print(f" {cond:<24} {cs['pages']} pages, mean {cs['mean_chars_page']} chars/page (grid cap {cs['grid_chars_page']})") + print( + f" {cond:<24} {cs['pages']} pages, mean {cs['mean_chars_page']} chars/page (grid cap {cs['grid_chars_page']})" + ) if args.render_only: for cond, u, ctx in tasks: @@ -418,9 +462,13 @@ def main() -> None: sub = [r for r in records if r["length"] == length and r["cond"] == cond] if not sub: continue - cells.append({"model": MODEL, "length": length, "condition": cond, **aggregate(sub)}) + cells.append( + {"model": MODEL, "length": length, "condition": cond, **aggregate(sub)} + ) (out_dir / "summary.json").write_text( - json.dumps({"args": vars(args), "capacity": capacity_stats, "cells": cells}, indent=1) + json.dumps( + {"args": vars(args), "capacity": capacity_stats, "cells": cells}, indent=1 + ) ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) diff --git a/packages/snapcompact/research/exp19_bestglm.py b/packages/snapcompact/research/exp19_bestglm.py index b8ef651ba..fe297d698 100644 --- a/packages/snapcompact/research/exp19_bestglm.py +++ b/packages/snapcompact/research/exp19_bestglm.py @@ -48,7 +48,9 @@ from run import CACHE, QA_CACHE, RESULTS, load_prompt, sha8 # noqa: E402 MODELS = {"z-ai/glm-4.6v": (0.30, 0.90)} LENGTHS = (150,) FONTS = { - "8on16": FontCfg("8on16", "8x13", 8, 16), # patch-aligned padded cell (exp01 pattern) + "8on16": FontCfg( + "8on16", "8x13", 8, 16 + ), # patch-aligned padded cell (exp01 pattern) "8x13": FontCfg("8x13", "8x13", 8, 13), # baseline pitch } # condition -> (kind, font key, palette variant) @@ -125,7 +127,9 @@ def pack_pages(paras: list[dict], col_w: int, max_lines: int) -> list[tuple[int, return pages -def _sentence_colors(lines: list[dict], palette: list) -> list[list[tuple[int, int, int]]]: +def _sentence_colors( + lines: list[dict], palette: list +) -> list[list[tuple[int, int, int]]]: """Per-line per-char glyph color cycling hue per sentence across the page.""" joined = "\n".join(ln["text"] for ln in lines) idx, out_idx = 0, [] @@ -141,7 +145,9 @@ def _sentence_colors(lines: list[dict], palette: list) -> list[list[tuple[int, i return colors -def render_doc(lines: list[dict], font: FontCfg, size: int, variant: str, cache: Path) -> Image.Image: +def render_doc( + lines: list[dict], font: FontCfg, size: int, variant: str, cache: Path +) -> Image.Image: """Two-column page: left column rows top-to-bottom, then right column. variant "dark-sent": black page, body glyphs in bright sentence hues, @@ -198,10 +204,20 @@ def save_png(png: Path, img_fn) -> None: tmp.replace(png) -def run_grid_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> list[dict]: +def run_grid_chunk( + model: str, cond: str, start: int, end: int, ctx: dict +) -> list[dict]: """One row-major-grid chunk: render via bdf.render, QA, score.""" - args, flow, paras, offsets, keys = ctx["args"], ctx["flow"], ctx["paras"], ctx["offsets"], ctx["keys"] - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + args, flow, paras, offsets, keys = ( + ctx["args"], + ctx["flow"], + ctx["paras"], + ctx["offsets"], + ctx["keys"], + ) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] _, font_key, variant = CONDITIONS[cond] @@ -220,13 +236,17 @@ def run_doc_page(model: str, cond: str, page: tuple[int, int], ctx: dict) -> lis i, j = page start = offsets[i] end = offsets[j - 1] + len(paras[j - 1]["ctx"]) - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] _, font_key, variant = CONDITIONS[cond] font = FONTS[font_key] lines = ctx["lines"][cond][page] - key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size)) + key = sha8( + cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size) + ) png = CACHE / f"exp19-doc-{font_key}-{variant}-{key}.png" save_png(png, lambda: render_doc(lines, font, args.size, variant, CACHE)) cols, rows, _ = capacity(font, args.size) @@ -254,7 +274,15 @@ def parse_answers(text: str, n: int) -> list[str]: return nums -def qa_and_score(model: str, cond: str, prompt: str, png: Path, questions: list[dict], start: int, ctx: dict) -> list[dict]: +def qa_and_score( + model: str, + cond: str, + prompt: str, + png: Path, + questions: list[dict], + start: int, + ctx: dict, +) -> list[dict]: args, keys = ctx["args"], ctx["keys"] q_block = "\n".join(f"{k + 1}. {q['q']}" for k, q in enumerate(questions)) messages = [ @@ -268,11 +296,19 @@ def qa_and_score(model: str, cond: str, prompt: str, png: Path, questions: list[ } ] qa = cached( - model, "exp19-qa", {"messages": messages, "effort": args.effort}, + model, + "exp19-qa", + {"messages": messages, "effort": args.effort}, lambda: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, model, messages, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + model, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -305,8 +341,13 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out return { "n": n, @@ -333,7 +374,11 @@ def main() -> None: ap.add_argument("--max-tokens", type=int, default=32768) ap.add_argument("--effort", default=None) ap.add_argument("--fresh", action="store_true") - ap.add_argument("--render-only", action="store_true", help="render sample pages + capacity stats, no API") + ap.add_argument( + "--render-only", + action="store_true", + help="render sample pages + capacity stats, no API", + ) ap.add_argument("--env", default="~/.env") args = ap.parse_args() @@ -370,8 +415,12 @@ def main() -> None: cols, rows, _cap = capacity(font, args.size) col_w = (cols - GUTTER) // 2 pages = pack_pages(paras, col_w, 2 * rows) - page_lines[cond] = {pg: layout_page(paras[pg[0] : pg[1]], col_w) for pg in pages} - chars = [offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] for i, j in pages] + page_lines[cond] = { + pg: layout_page(paras[pg[0] : pg[1]], col_w) for pg in pages + } + chars = [ + offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] for i, j in pages + ] doc_stats[cond] = { "pages": len(pages), "mean_chars_page": round(sum(chars) / len(pages)), @@ -381,13 +430,19 @@ def main() -> None: capacity_stats[length] = { "corpus_chars": len(flow), "grid": { - fk: dict(zip(("cols", "rows", "chars"), capacity(FONTS[fk], args.size))) for fk in FONTS + fk: dict(zip(("cols", "rows", "chars"), capacity(FONTS[fk], args.size))) + for fk in FONTS }, "doc": doc_stats, } ctx = { - "args": args, "flow": flow, "paras": paras, "offsets": offsets, "keys": keys, - "length": length, "lines": page_lines, + "args": args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + "lines": page_lines, } for model in models: for cond in conditions: @@ -395,7 +450,15 @@ def main() -> None: if kind == "grid": budget = capacity(FONTS[font_key], args.size)[2] for start in range(0, len(flow), budget): - tasks.append(("grid", model, cond, (start, min(start + budget, len(flow))), ctx)) + tasks.append( + ( + "grid", + model, + cond, + (start, min(start + budget, len(flow))), + ctx, + ) + ) else: for pg in page_lines[cond]: tasks.append(("doc", model, cond, pg, ctx)) @@ -405,7 +468,9 @@ def main() -> None: for fk, g in st["grid"].items(): print(f" grid {fk}: {g['cols']}x{g['rows']} = {g['chars']} chars/page") for cond, d in st["doc"].items(): - print(f" {cond}: {d['pages']} pages, mean {d['mean_chars_page']} chars/page (2x{d['col_w']}w, {d['rows']} rows)") + print( + f" {cond}: {d['pages']} pages, mean {d['mean_chars_page']} chars/page (2x{d['col_w']}w, {d['rows']} rows)" + ) if args.render_only: for length in lengths: @@ -415,14 +480,35 @@ def main() -> None: if kind == "grid": budget = capacity(FONTS[font_key], args.size)[2] chunk_text = ctx["flow"][:budget] - png = CACHE / f"exp19-{font_key}-{variant}-{sha8(chunk_text, str(args.size))}.png" - save_png(png, lambda: render(chunk_text, FONTS[font_key], CACHE, args.size, variant)) + png = ( + CACHE + / f"exp19-{font_key}-{variant}-{sha8(chunk_text, str(args.size))}.png" + ) + save_png( + png, + lambda: render( + chunk_text, FONTS[font_key], CACHE, args.size, variant + ), + ) else: pg = next(iter(ctx["lines"][cond])) i, j = pg - key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in ctx["paras"][i:j]]), str(args.size)) + key = sha8( + cond, + json.dumps([(p["title"], p["ctx"]) for p in ctx["paras"][i:j]]), + str(args.size), + ) png = CACHE / f"exp19-doc-{font_key}-{variant}-{key}.png" - save_png(png, lambda: render_doc(ctx["lines"][cond][pg], FONTS[font_key], args.size, variant, CACHE)) + save_png( + png, + lambda: render_doc( + ctx["lines"][cond][pg], + FONTS[font_key], + args.size, + variant, + CACHE, + ), + ) print(f" sample: {png}") return @@ -431,7 +517,9 @@ def main() -> None: done = 0 with ThreadPoolExecutor(args.workers) as pool: futures = [ - pool.submit(run_grid_chunk, m, c, u[0], u[1], ctx) if kind == "grid" else pool.submit(run_doc_page, m, c, u, ctx) + pool.submit(run_grid_chunk, m, c, u[0], u[1], ctx) + if kind == "grid" + else pool.submit(run_doc_page, m, c, u, ctx) for kind, m, c, u, ctx in tasks ] for fut in futures: @@ -444,7 +532,9 @@ def main() -> None: rec_path = out_dir / "records.jsonl" if rec_path.exists(): old = [json.loads(ln) for ln in rec_path.read_text().splitlines() if ln.strip()] - records = [r for r in old if (r["model"], r["length"], r["cond"]) not in ran] + records + records = [ + r for r in old if (r["model"], r["length"], r["cond"]) not in ran + ] + records with rec_path.open("w") as fh: for r in records: fh.write(json.dumps(r) + "\n") @@ -453,12 +543,27 @@ def main() -> None: for model in sorted({r["model"] for r in records}): for length in sorted({r["length"] for r in records}): for cond in CONDITIONS: - sub = [r for r in records if r["model"] == model and r["length"] == length and r["cond"] == cond] + sub = [ + r + for r in records + if r["model"] == model + and r["length"] == length + and r["cond"] == cond + ] if not sub: continue - cells.append({"model": model, "length": length, "condition": cond, **aggregate(sub, *MODELS[model])}) + cells.append( + { + "model": model, + "length": length, + "condition": cond, + **aggregate(sub, *MODELS[model]), + } + ) (out_dir / "summary.json").write_text( - json.dumps({"args": vars(args), "capacity": capacity_stats, "cells": cells}, indent=1) + json.dumps( + {"args": vars(args), "capacity": capacity_stats, "cells": cells}, indent=1 + ) ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) diff --git a/packages/snapcompact/research/exp20_8x8u.py b/packages/snapcompact/research/exp20_8x8u.py index a23c063e5..206396907 100644 --- a/packages/snapcompact/research/exp20_8x8u.py +++ b/packages/snapcompact/research/exp20_8x8u.py @@ -50,13 +50,31 @@ _INK = (24, 24, 24) # model -> (cond, layout, variant, price_in, price_out, key_name) CONFIGS = { "gpt-5.5": ("img-doc-8x8u-bw", "doc", "bw", 2.0, 16.0, "openai"), - "google/gemini-3.5-flash": ("img-doc-8x8u-sent-dim", "doc", "sent-dim", 0.6, 4.0, "openrouter"), - "moonshotai/kimi-k2.6": ("img-doc-8x8u-sent-dim", "doc", "sent-dim", 0.68, 3.41, "openrouter"), + "google/gemini-3.5-flash": ( + "img-doc-8x8u-sent-dim", + "doc", + "sent-dim", + 0.6, + 4.0, + "openrouter", + ), + "moonshotai/kimi-k2.6": ( + "img-doc-8x8u-sent-dim", + "doc", + "sent-dim", + 0.68, + 3.41, + "openrouter", + ), "z-ai/glm-4.6v": ("img-doc-8x8u-sent", "doc", "sent", 0.30, 0.90, "openrouter"), "claude-fable-5": ("img-8x8u-dim", "grid", "dim", 10.0, 50.0, "anthropic"), "claude-opus-4-8": ("img-8x8u-bw", "grid", "bw", 15.0, 75.0, "anthropic"), } -KEY_ENV = {"openai": "OPENAI_API_KEY", "openrouter": "OPENROUTER_API_KEY", "anthropic": "ANTHROPIC_API_KEY"} +KEY_ENV = { + "openai": "OPENAI_API_KEY", + "openrouter": "OPENROUTER_API_KEY", + "anthropic": "ANTHROPIC_API_KEY", +} def slug(model: str) -> str: @@ -198,7 +216,9 @@ def render_doc(lines: list[dict], size: int, variant: str, cache: Path) -> Image def atomic_save(img: Image.Image, png: Path) -> None: - tmp = png.with_suffix(f".{os.getpid()}.tmp.png") # pid-unique: parallel models share sent-dim PNGs + tmp = png.with_suffix( + f".{os.getpid()}.tmp.png" + ) # pid-unique: parallel models share sent-dim PNGs img.save(tmp) tmp.replace(png) @@ -216,24 +236,53 @@ def parse_answers(text: str, n: int) -> list[str]: return nums -def qa_unit(model: str, cond: str, prompt: str, png: Path, questions: list[dict], length: int, start: int, ctx: dict) -> list[dict]: +def qa_unit( + model: str, + cond: str, + prompt: str, + png: Path, + questions: list[dict], + length: int, + start: int, + ctx: dict, +) -> list[dict]: args, keys = ctx["args"], ctx["keys"] q_block = "\n".join(f"{k + 1}. {q['q']}" for k, q in enumerate(questions)) - messages = [{"role": "user", "content": [{"text": prompt}, {"image_path": png}, {"text": q_block}]}] + messages = [ + { + "role": "user", + "content": [{"text": prompt}, {"image_path": png}, {"text": q_block}], + } + ] qa = cached( - model, {"messages": messages, "effort": None}, - lambda: dict(zip(("text", "usage", "stop"), llm_complete(keys, model, messages, max_tokens=args.max_tokens))), + model, + {"messages": messages, "effort": None}, + lambda: dict( + zip( + ("text", "usage", "stop"), + llm_complete(keys, model, messages, max_tokens=args.max_tokens), + ) + ), args.fresh, ) answers = parse_answers(qa["text"], len(questions)) records = [] for q, a in zip(questions, answers): - records.append({ - "model": model, "length": length, "cond": cond, "chunk": start, - "pos_rel": q["pos_rel"], "q": q["q"], "answer": a, "golds": q["golds"], - "em": squad.exact_match(a, q["golds"]), "f1": squad.f1(a, q["golds"]), - "abstained": "unreadable" in a.lower(), - }) + records.append( + { + "model": model, + "length": length, + "cond": cond, + "chunk": start, + "pos_rel": q["pos_rel"], + "q": q["q"], + "answer": a, + "golds": q["golds"], + "em": squad.exact_match(a, q["golds"]), + "f1": squad.f1(a, q["golds"]), + "abstained": "unreadable" in a.lower(), + } + ) records[0]["usage"] = [{"phase": "qa", **qa["usage"]}] return records @@ -254,13 +303,23 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> list[di creads = sum(u.get("cache_r", 0) for u in usage) rsn = sum(u.get("reasoning", 0) for u in usage) cost = (tin + 0.1 * creads) * price_in / 1e6 + tout * price_out / 1e6 - out.append({ - "model": recs[0]["model"], "length": length, "condition": cond, "n": n, - "em": round(sum(r["em"] for r in recs) / n, 4), "f1": round(mean, 4), - "f1_se": round((var / n) ** 0.5, 4), "abstained": sum(r["abstained"] for r in recs), - "tok_in": tin, "tok_out": tout, "tok_cache_r": creads, "tok_reasoning": rsn, - "cost_usd": round(cost, 4), - }) + out.append( + { + "model": recs[0]["model"], + "length": length, + "condition": cond, + "n": n, + "em": round(sum(r["em"] for r in recs) / n, 4), + "f1": round(mean, 4), + "f1_se": round((var / n) ** 0.5, 4), + "abstained": sum(r["abstained"] for r in recs), + "tok_in": tin, + "tok_out": tout, + "tok_cache_r": creads, + "tok_reasoning": rsn, + "cost_usd": round(cost, 4), + } + ) return out @@ -286,10 +345,17 @@ def main() -> None: cols, rows, grid_cap = capacity(FONT, args.size) col_w = (cols - GUTTER) // 2 max_lines = 2 * rows - print(f"{args.model}: {cond} ({layout}/{variant}); 8x8u grid {cols}x{rows}={grid_cap}, " - f"doc 2x{col_w}+g{GUTTER}, {max_lines} slots", flush=True) + print( + f"{args.model}: {cond} ({layout}/{variant}); 8x8u grid {cols}x{rows}={grid_cap}, " + f"doc 2x{col_w}+g{GUTTER}, {max_lines} slots", + flush=True, + ) - keys = {} if args.render_only else {key_name: load_env_key(KEY_ENV[key_name], args.env)} + keys = ( + {} + if args.render_only + else {key_name: load_env_key(KEY_ENV[key_name], args.env)} + ) all_paras = squad.load_paragraphs(CACHE) tasks = [] cap_stats = {} @@ -299,35 +365,63 @@ def main() -> None: ctx = {"args": args, "keys": keys} if layout == "doc": pages = pack_pages(paras, col_w, max_lines) - page_chars = [offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] for i, j in pages] - cap_stats[length] = {"pages": len(pages), "mean_chars_page": round(sum(page_chars) / len(pages))} + page_chars = [ + offsets[j - 1] + len(paras[j - 1]["ctx"]) - offsets[i] for i, j in pages + ] + cap_stats[length] = { + "pages": len(pages), + "mean_chars_page": round(sum(page_chars) / len(pages)), + } prompt = load_prompt("exp04-qa-image.md").format(col_w=col_w, rows=rows) for i, j in pages: start = offsets[i] end = offsets[j - 1] + len(paras[j - 1]["ctx"]) - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: continue lines = layout_page(paras[i:j], col_w) - key = sha8(cond, json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), str(args.size)) + key = sha8( + cond, + json.dumps([(p["title"], p["ctx"]) for p in paras[i:j]]), + str(args.size), + ) png = CACHE / f"{EXP}-doc-{variant}-{key}.png" if not png.exists() or png.stat().st_size == 0: atomic_save(render_doc(lines, args.size, variant, CACHE), png) - tasks.append((args.model, cond, prompt, png, questions, length, start, ctx)) + tasks.append( + (args.model, cond, prompt, png, questions, length, start, ctx) + ) else: - cap_stats[length] = {"pages": -(-len(flow) // grid_cap), "mean_chars_page": grid_cap} + cap_stats[length] = { + "pages": -(-len(flow) // grid_cap), + "mean_chars_page": grid_cap, + } prompt = load_prompt("qa-image.md").format(cols=cols, rows=rows) for start in range(0, len(flow), grid_cap): end = min(start + grid_cap, len(flow)) - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: continue - png = CACHE / f"{EXP}-grid-{variant}-{sha8(flow[start:end], str(args.size))}.png" + png = ( + CACHE + / f"{EXP}-grid-{variant}-{sha8(flow[start:end], str(args.size))}.png" + ) if not png.exists() or png.stat().st_size == 0: - atomic_save(render(flow[start:end], FONT, CACHE, args.size, variant), png) - tasks.append((args.model, cond, prompt, png, questions, length, start, ctx)) - print(f" len {length}: {cap_stats[length]['pages']} pages, " - f"mean {cap_stats[length]['mean_chars_page']} chars/page, corpus {len(flow)}", flush=True) + atomic_save( + render(flow[start:end], FONT, CACHE, args.size, variant), png + ) + tasks.append( + (args.model, cond, prompt, png, questions, length, start, ctx) + ) + print( + f" len {length}: {cap_stats[length]['pages']} pages, " + f"mean {cap_stats[length]['mean_chars_page']} chars/page, corpus {len(flow)}", + flush=True, + ) if args.render_only: print(f"sample: {tasks[0][3]}" if tasks else "no tasks") @@ -350,11 +444,18 @@ def main() -> None: fh.write(hdr + "\n") for c in cells: fh.write(",".join(str(c[k]) for k in hdr.split(",")) + "\n") - (OUT_DIR / f"summary-{s}.json").write_text(json.dumps({"args": vars(args), "capacity": cap_stats, "cells": cells}, indent=1)) + (OUT_DIR / f"summary-{s}.json").write_text( + json.dumps( + {"args": vars(args), "capacity": cap_stats, "cells": cells}, indent=1 + ) + ) for c in cells: - print(f"len {c['length']:<4} {c['condition']:<24} n={c['n']:<4} EM {c['em']:.3f} " - f"F1 {c['f1']:.3f} ±{c['f1_se']:.3f} ${c['cost_usd']:.3f} " - f"out={c['tok_out']} rsn={c['tok_reasoning']}", flush=True) + print( + f"len {c['length']:<4} {c['condition']:<24} n={c['n']:<4} EM {c['em']:.3f} " + f"F1 {c['f1']:.3f} ±{c['f1_se']:.3f} ${c['cost_usd']:.3f} " + f"out={c['tok_out']} rsn={c['tok_reasoning']}", + flush=True, + ) print(f"-> {OUT_DIR}/records-{s}.jsonl", flush=True) diff --git a/packages/snapcompact/research/exp21_braille.py b/packages/snapcompact/research/exp21_braille.py index d553c7be0..684dbbdc6 100644 --- a/packages/snapcompact/research/exp21_braille.py +++ b/packages/snapcompact/research/exp21_braille.py @@ -49,14 +49,43 @@ _BLACK = (0, 0, 0) # dots numbered 1-6: 1=top-left 2=mid-left 3=bottom-left 4=top-right 5=mid-right 6=bottom-right # bitmask: bit0=dot1 .. bit5=dot6 (matches Unicode U+2800 offsets) _L = { - "a": 0x01, "b": 0x03, "c": 0x09, "d": 0x19, "e": 0x11, "f": 0x0B, "g": 0x1B, - "h": 0x13, "i": 0x0A, "j": 0x1A, "k": 0x05, "l": 0x07, "m": 0x0D, "n": 0x1D, - "o": 0x15, "p": 0x0F, "q": 0x1F, "r": 0x17, "s": 0x0E, "t": 0x1E, "u": 0x25, - "v": 0x27, "w": 0x3A, "x": 0x2D, "y": 0x3D, "z": 0x35, + "a": 0x01, + "b": 0x03, + "c": 0x09, + "d": 0x19, + "e": 0x11, + "f": 0x0B, + "g": 0x1B, + "h": 0x13, + "i": 0x0A, + "j": 0x1A, + "k": 0x05, + "l": 0x07, + "m": 0x0D, + "n": 0x1D, + "o": 0x15, + "p": 0x0F, + "q": 0x1F, + "r": 0x17, + "s": 0x0E, + "t": 0x1E, + "u": 0x25, + "v": 0x27, + "w": 0x3A, + "x": 0x2D, + "y": 0x3D, + "z": 0x35, } _PUNCT = { - ".": 0x32, ",": 0x02, "'": 0x04, "-": 0x24, ":": 0x12, ";": 0x06, - "?": 0x26, "!": 0x16, " ": 0x00, + ".": 0x32, + ",": 0x02, + "'": 0x04, + "-": 0x24, + ":": 0x12, + ";": 0x06, + "?": 0x26, + "!": 0x16, + " ": 0x00, } _NUMSIGN = 0x3C # dots 3456 _DIGIT = {d: _L["abcdefghij"[i]] for i, d in enumerate("1234567890")} @@ -132,28 +161,60 @@ def cached(payload: object, fn, fresh: bool) -> dict: return out -def qa_unit(cond: str, prompt: str, png: Path, questions: list[dict], length: int, start: int, ctx: dict) -> list[dict]: +def qa_unit( + cond: str, + prompt: str, + png: Path, + questions: list[dict], + length: int, + start: int, + ctx: dict, +) -> list[dict]: args, keys = ctx["args"], ctx["keys"] q_block = "\n".join(f"{k + 1}. {q['q']}" for k, q in enumerate(questions)) - messages = [{"role": "user", "content": [{"text": prompt}, {"image_path": png}, {"text": q_block}]}] + messages = [ + { + "role": "user", + "content": [{"text": prompt}, {"image_path": png}, {"text": q_block}], + } + ] payload = {"messages": messages} if args.effort: payload["effort"] = args.effort qa = cached( payload, - lambda: dict(zip(("text", "usage", "stop"), - llm_complete(keys, MODEL, messages, max_tokens=args.max_tokens, effort=args.effort))), + lambda: dict( + zip( + ("text", "usage", "stop"), + llm_complete( + keys, + MODEL, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), + ) + ), args.fresh, ) answers = squad.parse_numbered(qa["text"], len(questions)) records = [] for q, a in zip(questions, answers): - records.append({ - "model": MODEL, "length": length, "cond": cond, "chunk": start, - "pos_rel": q["pos_rel"], "q": q["q"], "answer": a, "golds": q["golds"], - "em": squad.exact_match(a, q["golds"]), "f1": squad.f1(a, q["golds"]), - "abstained": "unreadable" in a.lower(), - }) + records.append( + { + "model": MODEL, + "length": length, + "cond": cond, + "chunk": start, + "pos_rel": q["pos_rel"], + "q": q["q"], + "answer": a, + "golds": q["golds"], + "em": squad.exact_match(a, q["golds"]), + "f1": squad.f1(a, q["golds"]), + "abstained": "unreadable" in a.lower(), + } + ) records[0]["usage"] = [{"phase": "qa", **qa["usage"]}] return records @@ -174,20 +235,32 @@ def aggregate(records: list[dict]) -> list[dict]: creads = sum(u.get("cache_r", 0) for u in usage) rsn = sum(u.get("reasoning", 0) for u in usage) cost = (tin + 0.1 * creads) * PRICE_IN / 1e6 + tout * PRICE_OUT / 1e6 - out.append({ - "model": MODEL, "length": length, "condition": cond, "n": n, - "em": round(sum(r["em"] for r in recs) / n, 4), "f1": round(mean, 4), - "f1_se": round((var / n) ** 0.5, 4), "abstained": sum(r["abstained"] for r in recs), - "tok_in": tin, "tok_out": tout, "tok_cache_r": creads, "tok_reasoning": rsn, - "cost_usd": round(cost, 4), - }) + out.append( + { + "model": MODEL, + "length": length, + "condition": cond, + "n": n, + "em": round(sum(r["em"] for r in recs) / n, 4), + "f1": round(mean, 4), + "f1_se": round((var / n) ** 0.5, 4), + "abstained": sum(r["abstained"] for r in recs), + "tok_in": tin, + "tok_out": tout, + "tok_cache_r": creads, + "tok_reasoning": rsn, + "cost_usd": round(cost, 4), + } + ) return out def main() -> None: global MODEL, PRICE_IN, PRICE_OUT ap = argparse.ArgumentParser() - ap.add_argument("--model", default="google/gemini-3.5-flash", choices=sorted(MODELS)) + ap.add_argument( + "--model", default="google/gemini-3.5-flash", choices=sorted(MODELS) + ) ap.add_argument("--cells", default="5x7,7x10") ap.add_argument("--lengths", default="50,150") ap.add_argument("--qpc", type=int, default=30) @@ -216,7 +289,9 @@ def main() -> None: dpx, adv, pitch = CELLS[cell_name] cols, rows = args.size // adv, args.size // pitch cap = cols * rows - cond = f"img-braille-{cell_name}" + (f"+eff-{args.effort}" if args.effort else "") + cond = f"img-braille-{cell_name}" + ( + f"+eff-{args.effort}" if args.effort else "" + ) for length in (int(x) for x in args.lengths.split(",")): paras = all_paras[:length] flow, offsets = squad.build_flow(paras) @@ -227,16 +302,24 @@ def main() -> None: j = min(i + cap, len(cells)) pages.append((i, j)) i = j - print(f"{cond} len {length}: {len(pages)} pages, {cap} cells/page " - f"({cols}x{rows}), {len(cells)} cells for {len(flow)} chars", flush=True) + print( + f"{cond} len {length}: {len(pages)} pages, {cap} cells/page " + f"({cols}x{rows}), {len(cells)} cells for {len(flow)} chars", + flush=True, + ) ctx = {"args": args, "keys": keys} for ci, cj in pages: start = origin[ci] end = origin[cj - 1] + 1 - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: continue - png = CACHE / f"{EXP}-{cell_name}-{sha8(flow[start:end], cell_name, str(args.size))}.png" + png = ( + CACHE + / f"{EXP}-{cell_name}-{sha8(flow[start:end], cell_name, str(args.size))}.png" + ) if not png.exists() or png.stat().st_size == 0: atomic_save(render_braille(cells[ci:cj], cell_name, args.size), png) prompt = prompt_tpl.format(cols=cols, rows=rows) @@ -264,11 +347,16 @@ def main() -> None: fh.write(hdr + "\n") for c in cells_out: fh.write(",".join(str(c[k]) for k in hdr.split(",")) + "\n") - (OUT_DIR / f"summary-{slug}.json").write_text(json.dumps({"args": vars(args), "cells": cells_out}, indent=1)) + (OUT_DIR / f"summary-{slug}.json").write_text( + json.dumps({"args": vars(args), "cells": cells_out}, indent=1) + ) for c in cells_out: - print(f"len {c['length']:<4} {c['condition']:<20} n={c['n']:<4} EM {c['em']:.3f} " - f"F1 {c['f1']:.3f} ±{c['f1_se']:.3f} ${c['cost_usd']:.3f} " - f"out={c['tok_out']} rsn={c['tok_reasoning']}", flush=True) + print( + f"len {c['length']:<4} {c['condition']:<20} n={c['n']:<4} EM {c['em']:.3f} " + f"F1 {c['f1']:.3f} ±{c['f1_se']:.3f} ${c['cost_usd']:.3f} " + f"out={c['tok_out']} rsn={c['tok_reasoning']}", + flush=True, + ) print(f"-> {OUT_DIR}/matrix-{slug}.csv", flush=True) diff --git a/packages/snapcompact/research/exp22_ttf6pt.py b/packages/snapcompact/research/exp22_ttf6pt.py index 6d95db910..c720d2784 100644 --- a/packages/snapcompact/research/exp22_ttf6pt.py +++ b/packages/snapcompact/research/exp22_ttf6pt.py @@ -107,28 +107,60 @@ def cached(payload: object, fn, fresh: bool) -> dict: return out -def qa_unit(cond: str, prompt: str, png: Path, questions: list[dict], length: int, start: int, ctx: dict) -> list[dict]: +def qa_unit( + cond: str, + prompt: str, + png: Path, + questions: list[dict], + length: int, + start: int, + ctx: dict, +) -> list[dict]: args, keys = ctx["args"], ctx["keys"] q_block = "\n".join(f"{k + 1}. {q['q']}" for k, q in enumerate(questions)) - messages = [{"role": "user", "content": [{"text": prompt}, {"image_path": png}, {"text": q_block}]}] + messages = [ + { + "role": "user", + "content": [{"text": prompt}, {"image_path": png}, {"text": q_block}], + } + ] payload = {"messages": messages} if args.effort: payload["effort"] = args.effort qa = cached( payload, - lambda: dict(zip(("text", "usage", "stop"), - llm_complete(keys, MODEL, messages, max_tokens=args.max_tokens, effort=args.effort))), + lambda: dict( + zip( + ("text", "usage", "stop"), + llm_complete( + keys, + MODEL, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + ), + ) + ), args.fresh, ) answers = squad.parse_numbered(qa["text"], len(questions)) records = [] for q, a in zip(questions, answers): - records.append({ - "model": MODEL, "length": length, "cond": cond, "chunk": start, - "pos_rel": q["pos_rel"], "q": q["q"], "answer": a, "golds": q["golds"], - "em": squad.exact_match(a, q["golds"]), "f1": squad.f1(a, q["golds"]), - "abstained": "unreadable" in a.lower(), - }) + records.append( + { + "model": MODEL, + "length": length, + "cond": cond, + "chunk": start, + "pos_rel": q["pos_rel"], + "q": q["q"], + "answer": a, + "golds": q["golds"], + "em": squad.exact_match(a, q["golds"]), + "f1": squad.f1(a, q["golds"]), + "abstained": "unreadable" in a.lower(), + } + ) records[0]["usage"] = [{"phase": "qa", **qa["usage"]}] return records @@ -149,20 +181,32 @@ def aggregate(records: list[dict]) -> list[dict]: creads = sum(u.get("cache_r", 0) for u in usage) rsn = sum(u.get("reasoning", 0) for u in usage) cost = (tin + 0.1 * creads) * PRICE_IN / 1e6 + tout * PRICE_OUT / 1e6 - out.append({ - "model": MODEL, "length": length, "condition": cond, "n": n, - "em": round(sum(r["em"] for r in recs) / n, 4), "f1": round(mean, 4), - "f1_se": round((var / n) ** 0.5, 4), "abstained": sum(r["abstained"] for r in recs), - "tok_in": tin, "tok_out": tout, "tok_cache_r": creads, "tok_reasoning": rsn, - "cost_usd": round(cost, 4), - }) + out.append( + { + "model": MODEL, + "length": length, + "condition": cond, + "n": n, + "em": round(sum(r["em"] for r in recs) / n, 4), + "f1": round(mean, 4), + "f1_se": round((var / n) ** 0.5, 4), + "abstained": sum(r["abstained"] for r in recs), + "tok_in": tin, + "tok_out": tout, + "tok_cache_r": creads, + "tok_reasoning": rsn, + "cost_usd": round(cost, 4), + } + ) return out def main() -> None: global MODEL, PRICE_IN, PRICE_OUT ap = argparse.ArgumentParser() - ap.add_argument("--model", default="google/gemini-3.5-flash", choices=sorted(MODELS)) + ap.add_argument( + "--model", default="google/gemini-3.5-flash", choices=sorted(MODELS) + ) ap.add_argument("--ems", default="6,8") ap.add_argument("--lengths", default="50,150") ap.add_argument("--qpc", type=int, default=30) @@ -191,7 +235,10 @@ def main() -> None: adv, pitch, cols, rows = metrics(em) cap = cols * rows cond = f"img-ttf{em}-bw" + (f"+eff-{args.effort}" if args.effort else "") - print(f"{cond}: adv {adv:.2f}px pitch {pitch}px -> {cols}x{rows} = {cap} chars/page", flush=True) + print( + f"{cond}: adv {adv:.2f}px pitch {pitch}px -> {cols}x{rows} = {cap} chars/page", + flush=True, + ) for length in (int(x) for x in args.lengths.split(",")): paras = all_paras[:length] flow, offsets = squad.build_flow(paras) @@ -200,10 +247,15 @@ def main() -> None: ctx = {"args": args, "keys": keys} for start in range(0, len(flow), cap): end = min(start + cap, len(flow)) - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: continue - png = CACHE / f"{EXP}-ttf{em}-{sha8(flow[start:end], str(em), str(args.size))}.png" + png = ( + CACHE + / f"{EXP}-ttf{em}-{sha8(flow[start:end], str(em), str(args.size))}.png" + ) if not png.exists() or png.stat().st_size == 0: atomic_save(render_ttf(flow[start:end], em, args.size), png) prompt = prompt_tpl.format(cols=cols, rows=rows) @@ -231,11 +283,16 @@ def main() -> None: fh.write(hdr + "\n") for c in cells_out: fh.write(",".join(str(c[k]) for k in hdr.split(",")) + "\n") - (OUT_DIR / f"summary-{slug}.json").write_text(json.dumps({"args": vars(args), "cells": cells_out}, indent=1)) + (OUT_DIR / f"summary-{slug}.json").write_text( + json.dumps({"args": vars(args), "cells": cells_out}, indent=1) + ) for c in cells_out: - print(f"len {c['length']:<4} {c['condition']:<20} n={c['n']:<4} EM {c['em']:.3f} " - f"F1 {c['f1']:.3f} ±{c['f1_se']:.3f} ${c['cost_usd']:.3f} " - f"out={c['tok_out']} rsn={c['tok_reasoning']}", flush=True) + print( + f"len {c['length']:<4} {c['condition']:<20} n={c['n']:<4} EM {c['em']:.3f} " + f"F1 {c['f1']:.3f} ±{c['f1_se']:.3f} ${c['cost_usd']:.3f} " + f"out={c['tok_out']} rsn={c['tok_reasoning']}", + flush=True, + ) print(f"-> {OUT_DIR}/matrix-{slug}.csv", flush=True) diff --git a/packages/snapcompact/research/final.py b/packages/snapcompact/research/final.py index a3b5cd509..7e75032c9 100644 --- a/packages/snapcompact/research/final.py +++ b/packages/snapcompact/research/final.py @@ -38,7 +38,16 @@ sys.path.insert(0, str(HERE)) import squad # noqa: E402 from bdf import capacity, render # noqa: E402 from providers import is_openai, llm_complete, load_env_key, openai_compact # noqa: E402 -from run import CACHE, FONTS, QA_CACHE, RESULTS, TEXT_CHUNK, agent_prompt, load_prompt, sha8 # noqa: E402 +from run import ( + CACHE, + FONTS, + QA_CACHE, + RESULTS, + TEXT_CHUNK, + agent_prompt, + load_prompt, + sha8, +) # noqa: E402 # (family display, $/M input, $/M output). Cached reads bill at 0.1x input, # Anthropic cache writes at 1.25x. Edit prices here; `--report` recomputes. @@ -53,7 +62,15 @@ MODELS = { "z-ai/glm-4.6v": (0.30, 0.90), } LENGTHS = (50, 150, 250) -CONDITIONS = ("text", "handoff", "compact", "img-6x10-sent", "img-6x10-bw", "img-5x8-sent", "img-5x8-bw") +CONDITIONS = ( + "text", + "handoff", + "compact", + "img-6x10-sent", + "img-6x10-bw", + "img-5x8-sent", + "img-5x8-bw", +) ACK = "Noted. I have read the passages and will keep them in mind." @@ -91,15 +108,30 @@ def chunk_budget(cond: str, size: int) -> int: def session_frame(chunk_text: str) -> list[dict]: return [ - {"role": "user", "content": [{"text": load_prompt("session-frame.md").format(context=chunk_text)}]}, + { + "role": "user", + "content": [ + {"text": load_prompt("session-frame.md").format(context=chunk_text)} + ], + }, {"role": "assistant", "content": [{"text": ACK}]}, ] -def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> list[dict]: +def run_cell_chunk( + model: str, cond: str, start: int, end: int, ctx: dict +) -> list[dict]: """One (model, condition, chunk) unit: build carrier, QA, score.""" - args, flow, paras, offsets, keys = ctx["args"], ctx["flow"], ctx["paras"], ctx["offsets"], ctx["keys"] - questions = squad.sample_chunk_questions(paras, offsets, start, end, args.qpc, args.seed) + args, flow, paras, offsets, keys = ( + ctx["args"], + ctx["flow"], + ctx["paras"], + ctx["offsets"], + ctx["keys"], + ) + questions = squad.sample_chunk_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if not questions: return [] chunk_text = flow[start:end] @@ -117,11 +149,15 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li png = CACHE / f"img-{tag}-{sha8(chunk_text, str(args.size), *salt)}.png" if not png.exists() or png.stat().st_size == 0: tmp = png.with_suffix(f".{uuid.uuid4().hex[:8]}.tmp.png") - render(chunk_text, FONTS[font], CACHE, args.size, variant, columns=columns).save(tmp) + render( + chunk_text, FONTS[font], CACHE, args.size, variant, columns=columns + ).save(tmp) tmp.replace(png) cols, rows, _ = capacity(FONTS[font], args.size, columns) preamble = ( - load_prompt("qa-image-cols.md").format(cols=cols, rows=rows, columns=columns) + load_prompt("qa-image-cols.md").format( + cols=cols, rows=rows, columns=columns + ) if columns > 1 else load_prompt("qa-image.md").format(cols=cols, rows=rows) ) @@ -137,25 +173,53 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li ] elif cond == "compact" and is_openai(model): comp = cached( - model, "remote-compact", {"chunk": chunk_text}, - lambda: dict(zip(("items", "usage"), openai_compact(keys["openai"], model, session_frame(chunk_text)))), + model, + "remote-compact", + {"chunk": chunk_text}, + lambda: dict( + zip( + ("items", "usage"), + openai_compact(keys["openai"], model, session_frame(chunk_text)), + ) + ), args.fresh, ) usage_rows.append(("compact", comp["usage"])) extra_items = comp["items"] messages = [ - {"role": "user", "content": [{"text": load_prompt("qa-remote-compact.md").format(questions=q_block)}]} + { + "role": "user", + "content": [ + { + "text": load_prompt("qa-remote-compact.md").format( + questions=q_block + ) + } + ], + } ] elif cond in ("compact", "handoff"): - prompt_file = {"compact": "compaction-summary.md", "handoff": "handoff-document.md"}[cond] + prompt_file = { + "compact": "compaction-summary.md", + "handoff": "handoff-document.md", + }[cond] gen = cached( - model, f"summary-{cond}", {"chunk": chunk_text}, + model, + f"summary-{cond}", + {"chunk": chunk_text}, lambda: dict( zip( ("text", "usage", "stop"), llm_complete( - keys, model, - session_frame(chunk_text) + [{"role": "user", "content": [{"text": agent_prompt(prompt_file)}]}], + keys, + model, + session_frame(chunk_text) + + [ + { + "role": "user", + "content": [{"text": agent_prompt(prompt_file)}], + } + ], system=agent_prompt("summarization-system.md"), max_tokens=args.max_tokens, ), @@ -167,25 +231,37 @@ def run_cell_chunk(model: str, cond: str, start: int, end: int, ctx: dict) -> li messages = [ { "role": "user", - "content": [{"text": load_prompt("qa-text.md").format(context=gen["text"])}, {"text": q_block}], + "content": [ + {"text": load_prompt("qa-text.md").format(context=gen["text"])}, + {"text": q_block}, + ], } ] else: # text messages = [ { "role": "user", - "content": [{"text": load_prompt("qa-text.md").format(context=chunk_text)}, {"text": q_block}], + "content": [ + {"text": load_prompt("qa-text.md").format(context=chunk_text)}, + {"text": q_block}, + ], } ] qa = cached( - model, "qa", {"messages": messages, "extra": extra_items, "effort": args.effort}, + model, + "qa", + {"messages": messages, "extra": extra_items, "effort": args.effort}, lambda: dict( zip( ("text", "usage", "stop"), llm_complete( - keys, model, messages, - max_tokens=args.max_tokens, effort=args.effort, extra_input_items=extra_items, + keys, + model, + messages, + max_tokens=args.max_tokens, + effort=args.effort, + extra_input_items=extra_items, ), ) ), @@ -220,8 +296,13 @@ def aggregate(records: list[dict], price_in: float, price_out: float) -> dict: mean_f1 = sum(f1s) / n se = (sum((x - mean_f1) ** 2 for x in f1s) / (n * (n - 1))) ** 0.5 if n > 1 else 0.0 us = [u for r in records if "usage" in r for u in r["usage"]] - tok = {k: sum(u.get(k, 0) for u in us) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost_in = (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + tok = { + k: sum(u.get(k, 0) for u in us) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost_in = ( + (tok["in"] + 1.25 * tok["cache_w"] + 0.1 * tok["cache_r"]) / 1e6 * price_in + ) cost_out = tok["out"] / 1e6 * price_out return { "n": n, @@ -246,11 +327,15 @@ def main() -> None: ap.add_argument("--size", type=int, default=1568) ap.add_argument("--workers", type=int, default=6) ap.add_argument("--max-tokens", type=int, default=16384) - ap.add_argument("--effort", default=None, help="reasoning effort; None = provider default") + ap.add_argument( + "--effort", default=None, help="reasoning effort; None = provider default" + ) ap.add_argument("--fresh", action="store_true") ap.add_argument("--report", action="store_true", help="reprint from cache only") ap.add_argument("--env", default="~/.env") - ap.add_argument("--out", default="final", help="results subdirectory (isolate concurrent runs)") + ap.add_argument( + "--out", default="final", help="results subdirectory (isolate concurrent runs)" + ) args = ap.parse_args() CACHE.mkdir(exist_ok=True) @@ -276,18 +361,31 @@ def main() -> None: for length in lengths: paras = all_paras[:length] flow, offsets = squad.build_flow(paras) - ctx = {"args": args, "flow": flow, "paras": paras, "offsets": offsets, "keys": keys, "length": length} + ctx = { + "args": args, + "flow": flow, + "paras": paras, + "offsets": offsets, + "keys": keys, + "length": length, + } for model in models: for cond in conditions: budget = chunk_budget(cond, args.size) for start in range(0, len(flow), budget): - tasks.append((model, cond, start, min(start + budget, len(flow)), ctx)) - print(f"grid: {len(models)} models x {len(lengths)} lengths x {len(conditions)} conditions = {len(tasks)} chunk tasks") + tasks.append( + (model, cond, start, min(start + budget, len(flow)), ctx) + ) + print( + f"grid: {len(models)} models x {len(lengths)} lengths x {len(conditions)} conditions = {len(tasks)} chunk tasks" + ) records: list[dict] = [] done = 0 with ThreadPoolExecutor(args.workers) as pool: - futures = [pool.submit(run_cell_chunk, m, c, s, e, ctx) for m, c, s, e, ctx in tasks] + futures = [ + pool.submit(run_cell_chunk, m, c, s, e, ctx) for m, c, s, e, ctx in tasks + ] for fut in futures: records.extend(fut.result()) done += 1 @@ -302,11 +400,26 @@ def main() -> None: for model in models: for length in lengths: for cond in conditions: - sub = [r for r in records if r["model"] == model and r["length"] == length and r["cond"] == cond] + sub = [ + r + for r in records + if r["model"] == model + and r["length"] == length + and r["cond"] == cond + ] if not sub: continue - cells.append({"model": model, "length": length, "condition": cond, **aggregate(sub, *MODELS[model])}) - (out_dir / "summary.json").write_text(json.dumps({"args": vars(args), "cells": cells}, indent=1)) + cells.append( + { + "model": model, + "length": length, + "condition": cond, + **aggregate(sub, *MODELS[model]), + } + ) + (out_dir / "summary.json").write_text( + json.dumps({"args": vars(args), "cells": cells}, indent=1) + ) with (out_dir / "matrix.csv").open("w", newline="") as fh: writer = csv.DictWriter(fh, fieldnames=list(cells[0].keys())) writer.writeheader() @@ -319,7 +432,16 @@ def main() -> None: for cond in conditions: row = f"{cond:<15}" for model in models: - cell = next((c for c in cells if c["model"] == model and c["length"] == length and c["condition"] == cond), None) + cell = next( + ( + c + for c in cells + if c["model"] == model + and c["length"] == length + and c["condition"] == cond + ), + None, + ) row += ( f"{cell['f1']:>10.3f} {cell['cost_in_usd']:>5.2f} {cell['cost_out_usd']:>5.2f}" if cell diff --git a/packages/snapcompact/research/mono.py b/packages/snapcompact/research/mono.py index db8071e63..9222da1ea 100644 --- a/packages/snapcompact/research/mono.py +++ b/packages/snapcompact/research/mono.py @@ -31,7 +31,9 @@ def build_content(cond: str, flow: str, size: int) -> tuple[list[dict], int]: img = parse_img_condition(cond) if not img: assert cond == "text", f"unsupported mono condition {cond!r}" - return [{"text": load_prompt("qa-text.md").format(context=flow), "cache": True}], 0 + return [ + {"text": load_prompt("qa-text.md").format(context=flow), "cache": True} + ], 0 font, variant, columns = img cfg = FONTS[font] cols, rows, cap = capacity(cfg, size, columns) @@ -46,14 +48,20 @@ def build_content(cond: str, flow: str, size: int) -> tuple[list[dict], int]: render(chunk, cfg, CACHE, size, variant, columns=columns).save(tmp) tmp.replace(png) pngs.append(png) - preamble = load_prompt("qa-image-multi.md").format(k=len(pngs), cols=cols, rows=rows) + preamble = load_prompt("qa-image-multi.md").format( + k=len(pngs), cols=cols, rows=rows + ) if cfg.repeat > 1: preamble += ( f"\nNote: every text line is rendered {cfg.repeat} times consecutively - first on the plain " "background, then repeated on a pale highlight band. The copies show identical characters; " "cross-check between them when a glyph is hard to read, and do not treat copies as separate text." ) - blocks = [{"text": preamble}, *({"image_path": p} for p in pngs), {"text": "End of images.", "cache": True}] + blocks = [ + {"text": preamble}, + *({"image_path": p} for p in pngs), + {"text": "End of images.", "cache": True}, + ] return blocks, len(pngs) @@ -61,9 +69,21 @@ def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--model", default="gpt-5.5") ap.add_argument("--chars", type=int, default=800_000) - ap.add_argument("--conditions", default="text,img-6x10-sent,img-6x8s-sent,img-8x8u-sent") - ap.add_argument("--questions", type=int, default=50, help="total questions sampled across the flow") - ap.add_argument("--qpb", type=int, default=5, help="questions per API call (context re-sent, prefix-cached)") + ap.add_argument( + "--conditions", default="text,img-6x10-sent,img-6x8s-sent,img-8x8u-sent" + ) + ap.add_argument( + "--questions", + type=int, + default=50, + help="total questions sampled across the flow", + ) + ap.add_argument( + "--qpb", + type=int, + default=5, + help="questions per API call (context re-sent, prefix-cached)", + ) ap.add_argument("--seed", type=int, default=42) ap.add_argument("--size", type=int, default=1568) ap.add_argument("--max-tokens", type=int, default=32768) @@ -80,7 +100,9 @@ def main() -> None: } paras = squad.load_paragraphs(CACHE) flow, offsets = squad.build_flow(paras, args.chars) - questions = squad.sample_chunk_questions(paras, offsets, 0, len(flow), args.questions, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, 0, len(flow), args.questions, args.seed + ) price_in, price_out = MODELS[args.model] print( f"flow: {len(flow):,} chars (~{len(flow) // 4 // 1000}k text tokens), " @@ -98,11 +120,19 @@ def main() -> None: q_block = "\n".join(f"{i + 1}. {q['q']}" for i, q in enumerate(batch)) messages = [{"role": "user", "content": [*ctx_blocks, {"text": q_block}]}] qa = cached( - args.model, "qa-mono", {"messages": messages, "effort": args.effort}, + args.model, + "qa-mono", + {"messages": messages, "effort": args.effort}, lambda m=messages: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, args.model, m, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + args.model, + m, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -112,29 +142,51 @@ def main() -> None: stops.append(qa["stop"]) rows = [ { - "model": args.model, "cond": cond, "pos_rel": q["pos_rel"], "q": q["q"], - "answer": a, "golds": q["golds"], - "em": squad.exact_match(a, q["golds"]), "f1": squad.f1(a, q["golds"]), + "model": args.model, + "cond": cond, + "pos_rel": q["pos_rel"], + "q": q["q"], + "answer": a, + "golds": q["golds"], + "em": squad.exact_match(a, q["golds"]), + "f1": squad.f1(a, q["golds"]), "abstained": "unreadable" in a.lower(), } for q, a in zip(questions, answers) ] - records.extend({**r, "usage": usages} if i == 0 else r for i, r in enumerate(rows)) - u = {k: sum(x[k] for x in usages) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} + records.extend( + {**r, "usage": usages} if i == 0 else r for i, r in enumerate(rows) + ) + u = { + k: sum(x[k] for x in usages) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } stop = next((s for s in stops if s == "max_tokens"), stops[-1] if stops else "") - cost = (u["in"] + 1.25 * u["cache_w"] + 0.1 * u["cache_r"]) / 1e6 * price_in + u["out"] / 1e6 * price_out + cost = ( + u["in"] + 1.25 * u["cache_w"] + 0.1 * u["cache_r"] + ) / 1e6 * price_in + u["out"] / 1e6 * price_out quart = [] for lo, hi in ((0, 0.25), (0.25, 0.5), (0.5, 0.75), (0.75, 1.01)): qs = [r["f1"] for r in rows if lo <= r["pos_rel"] < hi] quart.append(sum(qs) / len(qs) if qs else float("nan")) table.append( { - "cond": cond, "n": len(rows), "imgs": n_imgs, + "cond": cond, + "n": len(rows), + "imgs": n_imgs, "em": sum(r["em"] for r in rows) / len(rows), "f1": sum(r["f1"] for r in rows) / len(rows), "abst": sum(r["abstained"] for r in rows), - "tok_in": u["in"], "tok_cached": u["cache_r"], "tok_out": u["out"], "reas": u["reasoning"], - "cost": cost, "stop": stop, "q1": quart[0], "q2": quart[1], "q3": quart[2], "q4": quart[3], + "tok_in": u["in"], + "tok_cached": u["cache_r"], + "tok_out": u["out"], + "reas": u["reasoning"], + "cost": cost, + "stop": stop, + "q1": quart[0], + "q2": quart[1], + "q3": quart[2], + "q4": quart[3], } ) t = table[-1] @@ -143,7 +195,10 @@ def main() -> None: f"in={t['tok_in']:>7} cached={t['tok_cached']:>7} out={t['tok_out']:>6} reas={t['reas']:>6} " f"${t['cost']:.2f} stop={t['stop']}" ) - print(f"{'':<18} F1 by position quartile: " + " ".join(f"q{i + 1}={v:.3f}" for i, v in enumerate(quart))) + print( + f"{'':<18} F1 by position quartile: " + + " ".join(f"q{i + 1}={v:.3f}" for i, v in enumerate(quart)) + ) (out_dir / "records.jsonl").write_text("\n".join(json.dumps(r) for r in records)) (out_dir / "summary.json").write_text(json.dumps(table, indent=1)) diff --git a/packages/snapcompact/research/mono_prod.py b/packages/snapcompact/research/mono_prod.py index d3e1a0d59..a5abd83ee 100644 --- a/packages/snapcompact/research/mono_prod.py +++ b/packages/snapcompact/research/mono_prod.py @@ -31,39 +31,136 @@ SIZE = 1568 # Production Shape payloads, keyed by variant name (geometry only; billing # fields are required by isShape but irrelevant to rendering). SHAPES = { - "doc-8on16-bw": {"font": "8x13", "cellWidth": 8, "cellHeight": 16, "stretch": False, "variant": "bw", - "columns": 2, "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 2900}, - "doc-8on16-sent": {"font": "8x13", "cellWidth": 8, "cellHeight": 16, "stretch": False, "variant": "sent", - "columns": 2, "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 2900}, - "doc-8on16-sent-dim": {"font": "8x13", "cellWidth": 8, "cellHeight": 16, "stretch": False, "variant": "sent", - "columns": 2, "stopwordDim": True, "lineRepeat": 1, "frameSize": SIZE, - "frameTokenEstimate": 2900}, - "8on16-bw": {"font": "8x13", "cellWidth": 8, "cellHeight": 16, "stretch": False, "variant": "bw", - "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 2900}, - "6x12-dim": {"font": "6x12", "cellWidth": 6, "cellHeight": 12, "variant": "bw", "stopwordDim": True, - "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 3300}, - "8x13-bw": {"font": "8x13", "cellWidth": 8, "cellHeight": 13, "variant": "bw", - "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 3300}, - "8x8r-bw": {"font": "8x8", "cellWidth": 8, "cellHeight": 8, "variant": "bw", - "lineRepeat": 2, "frameSize": SIZE, "frameTokenEstimate": 3300}, - "8x8r-sent": {"font": "8x8", "cellWidth": 8, "cellHeight": 8, "variant": "sent", - "lineRepeat": 2, "frameSize": SIZE, "frameTokenEstimate": 1100}, - "8x8u-bw": {"font": "8x8", "cellWidth": 8, "cellHeight": 8, "variant": "bw", - "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 3300}, - "8x8u-sent": {"font": "8x8", "cellWidth": 8, "cellHeight": 8, "variant": "sent", - "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 3300}, - "6x6u-sent": {"font": "8x8", "cellWidth": 6, "cellHeight": 6, "variant": "sent", - "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 3300}, + "doc-8on16-bw": { + "font": "8x13", + "cellWidth": 8, + "cellHeight": 16, + "stretch": False, + "variant": "bw", + "columns": 2, + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 2900, + }, + "doc-8on16-sent": { + "font": "8x13", + "cellWidth": 8, + "cellHeight": 16, + "stretch": False, + "variant": "sent", + "columns": 2, + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 2900, + }, + "doc-8on16-sent-dim": { + "font": "8x13", + "cellWidth": 8, + "cellHeight": 16, + "stretch": False, + "variant": "sent", + "columns": 2, + "stopwordDim": True, + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 2900, + }, + "8on16-bw": { + "font": "8x13", + "cellWidth": 8, + "cellHeight": 16, + "stretch": False, + "variant": "bw", + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 2900, + }, + "6x12-dim": { + "font": "6x12", + "cellWidth": 6, + "cellHeight": 12, + "variant": "bw", + "stopwordDim": True, + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 3300, + }, + "8x13-bw": { + "font": "8x13", + "cellWidth": 8, + "cellHeight": 13, + "variant": "bw", + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 3300, + }, + "8x8r-bw": { + "font": "8x8", + "cellWidth": 8, + "cellHeight": 8, + "variant": "bw", + "lineRepeat": 2, + "frameSize": SIZE, + "frameTokenEstimate": 3300, + }, + "8x8r-sent": { + "font": "8x8", + "cellWidth": 8, + "cellHeight": 8, + "variant": "sent", + "lineRepeat": 2, + "frameSize": SIZE, + "frameTokenEstimate": 1100, + }, + "8x8u-bw": { + "font": "8x8", + "cellWidth": 8, + "cellHeight": 8, + "variant": "bw", + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 3300, + }, + "8x8u-sent": { + "font": "8x8", + "cellWidth": 8, + "cellHeight": 8, + "variant": "sent", + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 3300, + }, + "6x6u-sent": { + "font": "8x8", + "cellWidth": 6, + "cellHeight": 6, + "variant": "sent", + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 3300, + }, } def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--shape", choices=sorted(SHAPES), help="named shape from the built-in table") - ap.add_argument("--shape-json", help="raw production Shape JSON (alternative to --shape)") + ap.add_argument( + "--shape", choices=sorted(SHAPES), help="named shape from the built-in table" + ) + ap.add_argument( + "--shape-json", help="raw production Shape JSON (alternative to --shape)" + ) ap.add_argument("--name", help="condition label; required with --shape-json") - ap.add_argument("--price-in", type=float, help="$/M input tokens; overrides the final.MODELS table") - ap.add_argument("--price-out", type=float, help="$/M output tokens; overrides the final.MODELS table") + ap.add_argument( + "--price-in", + type=float, + help="$/M input tokens; overrides the final.MODELS table", + ) + ap.add_argument( + "--price-out", + type=float, + help="$/M output tokens; overrides the final.MODELS table", + ) ap.add_argument("--model", default="gpt-5.5") ap.add_argument("--chars", type=int, default=800_000) ap.add_argument("--questions", type=int, default=50) @@ -82,10 +179,16 @@ def main() -> None: } paras = squad.load_paragraphs(CACHE) flow, offsets = squad.build_flow(paras, args.chars) - questions = squad.sample_chunk_questions(paras, offsets, 0, len(flow), args.questions, args.seed) + questions = squad.sample_chunk_questions( + paras, offsets, 0, len(flow), args.questions, args.seed + ) table = MODELS.get(args.model) - price_in = args.price_in if args.price_in is not None else (table or (None, None))[0] - price_out = args.price_out if args.price_out is not None else (table or (None, None))[1] + price_in = ( + args.price_in if args.price_in is not None else (table or (None, None))[0] + ) + price_out = ( + args.price_out if args.price_out is not None else (table or (None, None))[1] + ) if price_in is None or price_out is None: ap.error(f"model {args.model} not in final.MODELS; pass --price-in/--price-out") @@ -101,19 +204,33 @@ def main() -> None: size = shape["frameSize"] # Production frames (keyed by flow + shape so corpus changes re-render). - frame_dir = CACHE / f"prod-frames-{label}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + frame_dir = ( + CACHE / f"prod-frames-{label}-{sha8(flow, json.dumps(shape, sort_keys=True))}" + ) if not frame_dir.exists() or not any(frame_dir.iterdir()): flow_file = CACHE / f"prod-flow-{sha8(flow)}.txt" flow_file.write_text(flow) subprocess.run( - ["bun", str(HERE / "render_pages.ts"), str(flow_file), json.dumps(shape), str(frame_dir)], + [ + "bun", + str(HERE / "render_pages.ts"), + str(flow_file), + json.dumps(shape), + str(frame_dir), + ], check=True, ) pngs = sorted(frame_dir.glob("page-*.png")) repeat = shape.get("lineRepeat", 1) - cols = (size // shape["cellWidth"] - 3) // 2 if shape.get("columns") == 2 else size // shape["cellWidth"] + cols = ( + (size // shape["cellWidth"] - 3) // 2 + if shape.get("columns") == 2 + else size // shape["cellWidth"] + ) rows = size // shape["cellHeight"] // repeat - preamble = load_prompt("qa-image-multi.md").format(k=len(pngs), cols=cols, rows=rows) + preamble = load_prompt("qa-image-multi.md").format( + k=len(pngs), cols=cols, rows=rows + ) if shape.get("columns") == 2: preamble += ( "\nNote: each image lays text out as two word-wrapped newspaper columns separated by a gutter; " @@ -125,7 +242,11 @@ def main() -> None: "background, then repeated on a pale highlight band. The copies show identical characters; " "cross-check between them when a glyph is hard to read, and do not treat copies as separate text." ) - ctx_blocks = [{"text": preamble}, *({"image_path": p} for p in pngs), {"text": "End of images.", "cache": True}] + ctx_blocks = [ + {"text": preamble}, + *({"image_path": p} for p in pngs), + {"text": "End of images.", "cache": True}, + ] out_dir = RESULTS / f"mono-prod-{args.model.replace('/', '-')}-{label}" out_dir.mkdir(parents=True, exist_ok=True) @@ -135,11 +256,19 @@ def main() -> None: q_block = "\n".join(f"{i + 1}. {q['q']}" for i, q in enumerate(batch)) messages = [{"role": "user", "content": [*ctx_blocks, {"text": q_block}]}] qa = cached( - args.model, "qa-mono-prod", {"messages": messages, "effort": args.effort}, + args.model, + "qa-mono-prod", + {"messages": messages, "effort": args.effort}, lambda m=messages: dict( zip( ("text", "usage", "stop"), - llm_complete(keys, args.model, m, max_tokens=args.max_tokens, effort=args.effort), + llm_complete( + keys, + args.model, + m, + max_tokens=args.max_tokens, + effort=args.effort, + ), ) ), args.fresh, @@ -149,26 +278,48 @@ def main() -> None: stops.append(qa["stop"]) rows_out = [ { - "model": args.model, "cond": cond, "pos_rel": q["pos_rel"], "q": q["q"], "answer": a, - "golds": q["golds"], "em": squad.exact_match(a, q["golds"]), "f1": squad.f1(a, q["golds"]), + "model": args.model, + "cond": cond, + "pos_rel": q["pos_rel"], + "q": q["q"], + "answer": a, + "golds": q["golds"], + "em": squad.exact_match(a, q["golds"]), + "f1": squad.f1(a, q["golds"]), "abstained": "unreadable" in a.lower(), } for q, a in zip(questions, answers) ] - u = {k: sum(x[k] for x in usages) for k in ("in", "out", "cache_w", "cache_r", "reasoning")} - cost = (u["in"] + 1.25 * u["cache_w"] + 0.1 * u["cache_r"]) / 1e6 * price_in + u["out"] / 1e6 * price_out + u = { + k: sum(x[k] for x in usages) + for k in ("in", "out", "cache_w", "cache_r", "reasoning") + } + cost = (u["in"] + 1.25 * u["cache_w"] + 0.1 * u["cache_r"]) / 1e6 * price_in + u[ + "out" + ] / 1e6 * price_out quart = [] for lo, hi in ((0, 0.25), (0.25, 0.5), (0.5, 0.75), (0.75, 1.01)): qs = [r["f1"] for r in rows_out if lo <= r["pos_rel"] < hi] quart.append(sum(qs) / len(qs) if qs else float("nan")) summary = { - "cond": cond, "n": len(rows_out), "imgs": len(pngs), + "cond": cond, + "n": len(rows_out), + "imgs": len(pngs), "em": sum(r["em"] for r in rows_out) / len(rows_out), "f1": sum(r["f1"] for r in rows_out) / len(rows_out), "abst": sum(r["abstained"] for r in rows_out), - "tok_in": u["in"], "tok_cached": u["cache_r"], "tok_out": u["out"], "reas": u["reasoning"], - "cost": cost, "stop": next((s for s in stops if s == "max_tokens"), stops[-1] if stops else ""), - "q1": quart[0], "q2": quart[1], "q3": quart[2], "q4": quart[3], + "tok_in": u["in"], + "tok_cached": u["cache_r"], + "tok_out": u["out"], + "reas": u["reasoning"], + "cost": cost, + "stop": next( + (s for s in stops if s == "max_tokens"), stops[-1] if stops else "" + ), + "q1": quart[0], + "q2": quart[1], + "q3": quart[2], + "q4": quart[3], } (out_dir / "records.jsonl").write_text("\n".join(json.dumps(r) for r in rows_out)) (out_dir / "summary.json").write_text(json.dumps([summary], indent=1)) @@ -176,7 +327,9 @@ def main() -> None: f"{cond:<22} imgs={summary['imgs']:>2} f1={summary['f1']:.3f} em={summary['em']:.3f} " f"abst={summary['abst']} ${summary['cost']:.2f} stop={summary['stop']}" ) - print("F1 by quartile: " + " ".join(f"q{i + 1}={v:.3f}" for i, v in enumerate(quart))) + print( + "F1 by quartile: " + " ".join(f"q{i + 1}={v:.3f}" for i, v in enumerate(quart)) + ) if __name__ == "__main__": diff --git a/packages/snapcompact/research/parity_check.py b/packages/snapcompact/research/parity_check.py index 55ec8c796..91bcc6c0d 100644 --- a/packages/snapcompact/research/parity_check.py +++ b/packages/snapcompact/research/parity_check.py @@ -43,40 +43,82 @@ SHAPES = [ FontCfg("6x12", "6x12", 6, 12), "dim", 1, - {"font": "6x12", "cellWidth": 6, "cellHeight": 12, "variant": "bw", "stopwordDim": True, - "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 3300}, + { + "font": "6x12", + "cellWidth": 6, + "cellHeight": 12, + "variant": "bw", + "stopwordDim": True, + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 3300, + }, ), ( "8x13-bw", FontCfg("8x13", "8x13", 8, 13), "bw", 1, - {"font": "8x13", "cellWidth": 8, "cellHeight": 13, "variant": "bw", - "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 3300}, + { + "font": "8x13", + "cellWidth": 8, + "cellHeight": 13, + "variant": "bw", + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 3300, + }, ), ( "8on16-bw", FontCfg("8on16", "8x13", 8, 16), "bw", 1, - {"font": "8x13", "cellWidth": 8, "cellHeight": 16, "stretch": False, "variant": "bw", - "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 3300}, + { + "font": "8x13", + "cellWidth": 8, + "cellHeight": 16, + "stretch": False, + "variant": "bw", + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 3300, + }, ), ( "doc-8on16-bw", FontCfg("8on16", "8x13", 8, 16), "bw", 2, - {"font": "8x13", "cellWidth": 8, "cellHeight": 16, "stretch": False, "variant": "bw", "columns": 2, - "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 3300}, + { + "font": "8x13", + "cellWidth": 8, + "cellHeight": 16, + "stretch": False, + "variant": "bw", + "columns": 2, + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 3300, + }, ), ( "doc-8on16-sent-dim", FontCfg("8on16", "8x13", 8, 16), "sent-dim", 2, - {"font": "8x13", "cellWidth": 8, "cellHeight": 16, "stretch": False, "variant": "sent", "columns": 2, - "stopwordDim": True, "lineRepeat": 1, "frameSize": SIZE, "frameTokenEstimate": 3300}, + { + "font": "8x13", + "cellWidth": 8, + "cellHeight": 16, + "stretch": False, + "variant": "sent", + "columns": 2, + "stopwordDim": True, + "lineRepeat": 1, + "frameSize": SIZE, + "frameTokenEstimate": 3300, + }, ), ] @@ -162,7 +204,9 @@ def render_research(flow: str, cfg: FontCfg, variant: str, columns: int) -> Imag def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--keep", action="store_true", help="keep PNG pairs in .cache/parity") + ap.add_argument( + "--keep", action="store_true", help="keep PNG pairs in .cache/parity" + ) args = ap.parse_args() PARITY.mkdir(parents=True, exist_ok=True) @@ -179,7 +223,13 @@ def main() -> None: ref = render_research(flow, cfg, variant, columns) out_png = PARITY / f"{name}.prod.png" proc = subprocess.run( - ["bun", str(HERE / "parity_render.ts"), str(text_file), json.dumps(shape), str(out_png)], + [ + "bun", + str(HERE / "parity_render.ts"), + str(text_file), + json.dumps(shape), + str(out_png), + ], capture_output=True, text=True, ) @@ -192,7 +242,9 @@ def main() -> None: # reference renders the full square. Compare the printed region # pixel-exact and require everything below it to be blank. if got.width != ref.width or got.height > ref.height: - print(f"FAIL {name}: size {got.size} incompatible with reference {ref.size}") + print( + f"FAIL {name}: size {got.size} incompatible with reference {ref.size}" + ) failures += 1 continue rpx, gpx = ref.load(), got.load() diff --git a/packages/snapcompact/research/providers.py b/packages/snapcompact/research/providers.py index 270cefd5d..58f4529ff 100644 --- a/packages/snapcompact/research/providers.py +++ b/packages/snapcompact/research/providers.py @@ -37,7 +37,9 @@ def load_env_key(var: str, env_path: str = "~/.env") -> str: def _post(url: str, body: dict, headers: dict, retries: int = 4) -> dict: payload = json.dumps(body).encode() - req = urllib.request.Request(url, data=payload, headers={"content-type": "application/json", **headers}) + req = urllib.request.Request( + url, data=payload, headers={"content-type": "application/json", **headers} + ) for attempt in range(retries + 1): try: with urllib.request.urlopen(req, timeout=600) as resp: @@ -76,7 +78,11 @@ def _anthropic_blocks(blocks: list[dict]) -> list[dict]: else: item = { "type": "image", - "source": {"type": "base64", "media_type": "image/png", "data": _png_b64(b["image_path"])}, + "source": { + "type": "base64", + "media_type": "image/png", + "data": _png_b64(b["image_path"]), + }, } if b.get("cache"): item["cache_control"] = {"type": "ephemeral"} @@ -85,12 +91,20 @@ def _anthropic_blocks(blocks: list[dict]) -> list[dict]: def _anthropic_complete( - api_key: str, model: str, messages: list[dict], system: str | None, max_tokens: int, effort: str | None + api_key: str, + model: str, + messages: list[dict], + system: str | None, + max_tokens: int, + effort: str | None, ) -> tuple[str, dict, str]: body: dict = { "model": model, "max_tokens": max_tokens, - "messages": [{"role": m["role"], "content": _anthropic_blocks(m["content"])} for m in messages], + "messages": [ + {"role": m["role"], "content": _anthropic_blocks(m["content"])} + for m in messages + ], } if system: body["system"] = system @@ -163,23 +177,40 @@ def _openai_complete( extra_input_items: list[dict] | None = None, ) -> tuple[str, dict, str]: input_items: list[dict] = list(extra_input_items or []) - input_items += [{"role": m["role"], "content": _openai_content(m["content"], m["role"])} for m in messages] - body: dict = {"model": model, "input": input_items, "max_output_tokens": max_tokens, "store": False} + input_items += [ + {"role": m["role"], "content": _openai_content(m["content"], m["role"])} + for m in messages + ] + body: dict = { + "model": model, + "input": input_items, + "max_output_tokens": max_tokens, + "store": False, + } if system: body["instructions"] = system if effort: body["reasoning"] = {"effort": "high" if effort in ("xhigh", "max") else effort} out = _post(OPENAI_URL, body, {"authorization": f"Bearer {api_key}"}) status = out.get("status", "") - stop = "max_tokens" if (out.get("incomplete_details") or {}).get("reason") == "max_output_tokens" else status + stop = ( + "max_tokens" + if (out.get("incomplete_details") or {}).get("reason") == "max_output_tokens" + else status + ) return _openai_output_text(out), _openai_usage(out), stop -def openai_compact(api_key: str, model: str, messages: list[dict]) -> tuple[list[dict], dict]: +def openai_compact( + api_key: str, model: str, messages: list[dict] +) -> tuple[list[dict], dict]: """POST /responses/compact: returns (compacted output items, usage).""" body = { "model": model, - "input": [{"role": m["role"], "content": _openai_content(m["content"], m["role"])} for m in messages], + "input": [ + {"role": m["role"], "content": _openai_content(m["content"], m["role"])} + for m in messages + ], } out = _post(f"{OPENAI_URL}/compact", body, {"authorization": f"Bearer {api_key}"}) return out.get("output", []), _openai_usage(out) @@ -191,7 +222,12 @@ OPENROUTER_URL = "https://openrouter.ai/api/v1/chat/completions" def _openrouter_complete( - api_key: str, model: str, messages: list[dict], system: str | None, max_tokens: int, effort: str | None + api_key: str, + model: str, + messages: list[dict], + system: str | None, + max_tokens: int, + effort: str | None, ) -> tuple[str, dict, str]: def content(blocks: list[dict]) -> list[dict]: out = [] @@ -199,15 +235,26 @@ def _openrouter_complete( if "text" in b: out.append({"type": "text", "text": b["text"]}) else: - out.append({"type": "image_url", "image_url": {"url": f"data:image/png;base64,{_png_b64(b['image_path'])}"}}) + out.append( + { + "type": "image_url", + "image_url": { + "url": f"data:image/png;base64,{_png_b64(b['image_path'])}" + }, + } + ) return out - chat_messages = [{"role": m["role"], "content": content(m["content"])} for m in messages] + chat_messages = [ + {"role": m["role"], "content": content(m["content"])} for m in messages + ] if system: chat_messages.insert(0, {"role": "system", "content": system}) body: dict = {"model": model, "messages": chat_messages, "max_tokens": max_tokens} if effort == "none": - body["reasoning"] = {"enabled": False} # OpenRouter's disable switch; effort "none" is not a valid level + body["reasoning"] = { + "enabled": False + } # OpenRouter's disable switch; effort "none" is not a valid level elif effort: body["reasoning"] = {"effort": "high" if effort in ("xhigh", "max") else effort} out = _post(OPENROUTER_URL, body, {"authorization": f"Bearer {api_key}"}) @@ -217,13 +264,20 @@ def _openrouter_complete( text = "".join(p.get("text", "") for p in text if isinstance(p, dict)) u = out.get("usage", {}) usage = { - "in": u.get("prompt_tokens", 0) - (u.get("prompt_tokens_details") or {}).get("cached_tokens", 0), + "in": u.get("prompt_tokens", 0) + - (u.get("prompt_tokens_details") or {}).get("cached_tokens", 0), "out": u.get("completion_tokens", 0), "cache_w": 0, "cache_r": (u.get("prompt_tokens_details") or {}).get("cached_tokens", 0), - "reasoning": (u.get("completion_tokens_details") or {}).get("reasoning_tokens", 0), + "reasoning": (u.get("completion_tokens_details") or {}).get( + "reasoning_tokens", 0 + ), } - stop = "max_tokens" if choice.get("finish_reason") == "length" else (choice.get("finish_reason") or "") + stop = ( + "max_tokens" + if choice.get("finish_reason") == "length" + else (choice.get("finish_reason") or "") + ) return text, usage, stop @@ -250,12 +304,24 @@ def llm_complete( """Returns (text, normalized usage, stop). stop == "max_tokens" means truncated.""" if is_openrouter(model): if extra_input_items: - raise ValueError("extra_input_items is OpenAI-only (compacted window replay)") - return _openrouter_complete(api_keys["openrouter"], model, messages, system, max_tokens, effort) + raise ValueError( + "extra_input_items is OpenAI-only (compacted window replay)" + ) + return _openrouter_complete( + api_keys["openrouter"], model, messages, system, max_tokens, effort + ) if is_openai(model): return _openai_complete( - api_keys["openai"], model, messages, system, max_tokens, effort, extra_input_items + api_keys["openai"], + model, + messages, + system, + max_tokens, + effort, + extra_input_items, ) if extra_input_items: raise ValueError("extra_input_items is OpenAI-only (compacted window replay)") - return _anthropic_complete(api_keys["anthropic"], model, messages, system, max_tokens, effort) + return _anthropic_complete( + api_keys["anthropic"], model, messages, system, max_tokens, effort + ) diff --git a/packages/snapcompact/research/snapcompact_3d_activation_html.py b/packages/snapcompact/research/snapcompact_3d_activation_html.py index f29726405..ed736e44e 100644 --- a/packages/snapcompact/research/snapcompact_3d_activation_html.py +++ b/packages/snapcompact/research/snapcompact_3d_activation_html.py @@ -41,7 +41,15 @@ def image_data_uri(path: Path) -> str: return "data:image/png;base64," + base64.b64encode(path.read_bytes()).decode() -def add_surface(fig: go.Figure, z: np.ndarray, row: int, col: int, name: str, colorscale: str, showscale: bool = False) -> None: +def add_surface( + fig: go.Figure, + z: np.ndarray, + row: int, + col: int, + name: str, + colorscale: str, + showscale: bool = False, +) -> None: y = np.arange(z.shape[0]) x = np.arange(z.shape[1]) fig.add_trace( @@ -54,11 +62,23 @@ def add_surface(fig: go.Figure, z: np.ndarray, row: int, col: int, name: str, co cmin=0, cmax=1, showscale=showscale, - lighting={"ambient": 0.58, "diffuse": 0.72, "specular": 0.28, "roughness": 0.52}, - contours={ - "z": {"show": True, "usecolormap": True, "highlightcolor": "#fff0a8", "project_z": True}, + lighting={ + "ambient": 0.58, + "diffuse": 0.72, + "specular": 0.28, + "roughness": 0.52, }, - hovertemplate="layer %{y}
image bin %{x}
Δ %{z:.3f}" + name + "", + contours={ + "z": { + "show": True, + "usecolormap": True, + "highlightcolor": "#fff0a8", + "project_z": True, + }, + }, + hovertemplate="layer %{y}
image bin %{x}
Δ %{z:.3f}" + + name + + "", ), row=row, col=col, @@ -67,8 +87,12 @@ def add_surface(fig: go.Figure, z: np.ndarray, row: int, col: int, name: str, co def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--result-dir", default=str(HERE / "results" / "tensor-heatmap-paddleocr-q7")) - ap.add_argument("--out", default=str(HERE / "results" / "snapcompact-activation-terrain.html")) + ap.add_argument( + "--result-dir", default=str(HERE / "results" / "tensor-heatmap-paddleocr-q7") + ) + ap.add_argument( + "--out", default=str(HERE / "results" / "snapcompact-activation-terrain.html") + ) ap.add_argument("--bins", type=int, default=150) args = ap.parse_args() @@ -82,22 +106,48 @@ def main() -> None: fig = make_subplots( rows=2, cols=2, - specs=[[{"type": "surface"}, {"type": "surface"}], [{"type": "surface", "colspan": 2}, None]], + specs=[ + [{"type": "surface"}, {"type": "surface"}], + [{"type": "surface", "colspan": 2}, None], + ], horizontal_spacing=0.02, vertical_spacing=0.03, - subplot_titles=("Gold answer erased", "Random equal-size erase", "Answer / random residual scar"), + subplot_titles=( + "Gold answer erased", + "Random equal-size erase", + "Answer / random residual scar", + ), ) add_surface(fig, answer, 1, 1, "gold answer mask", "Magma") add_surface(fig, random, 1, 2, "random mask", "Viridis") add_surface(fig, ratio, 2, 1, "answer/random ratio", "Inferno", True) - camera = {"eye": {"x": 1.65, "y": -1.75, "z": 0.82}, "center": {"x": 0, "y": 0, "z": -0.08}} + camera = { + "eye": {"x": 1.65, "y": -1.75, "z": 0.82}, + "center": {"x": 0, "y": 0, "z": -0.08}, + } scene_common = { "bgcolor": "rgba(0,0,0,0)", "camera": camera, - "xaxis": {"title": "image-token bins", "gridcolor": "rgba(140,170,180,0.18)", "color": "#94a3aa", "zeroline": False}, - "yaxis": {"title": "decoder layer", "gridcolor": "rgba(140,170,180,0.18)", "color": "#94a3aa", "autorange": "reversed", "dtick": 4}, - "zaxis": {"title": "Δ hidden", "gridcolor": "rgba(140,170,180,0.18)", "color": "#94a3aa", "range": [0, 1]}, + "xaxis": { + "title": "image-token bins", + "gridcolor": "rgba(140,170,180,0.18)", + "color": "#94a3aa", + "zeroline": False, + }, + "yaxis": { + "title": "decoder layer", + "gridcolor": "rgba(140,170,180,0.18)", + "color": "#94a3aa", + "autorange": "reversed", + "dtick": 4, + }, + "zaxis": { + "title": "Δ hidden", + "gridcolor": "rgba(140,170,180,0.18)", + "color": "#94a3aa", + "range": [0, 1], + }, "aspectratio": {"x": 2.6, "y": 0.78, "z": 0.52}, } fig.update_layout( @@ -117,7 +167,11 @@ def main() -> None: q = summary["question"] original_uri = image_data_uri(result_dir / "images" / "original.png") masked_uri = image_data_uri(result_dir / "images" / "answer-mask.png") - graph_html = fig.to_html(full_html=False, include_plotlyjs="cdn", config={"displayModeBar": False, "responsive": True}) + graph_html = fig.to_html( + full_html=False, + include_plotlyjs="cdn", + config={"displayModeBar": False, "responsive": True}, + ) html = f""" @@ -168,10 +222,10 @@ def main() -> None:
ANSWER ERASED
question
-
{q['q']}
+
{q["q"]}
gold answer
-
{q['answer_text']}
-
{summary['layers']} layers
{summary['image_tokens']} image tokens
answer/random Δ = {summary['answer_over_random_delta']:.2f}×
+
{q["answer_text"]}
+
{summary["layers"]} layers
{summary["image_tokens"]} image tokens
answer/random Δ = {summary["answer_over_random_delta"]:.2f}×
interactive 3D residual terraingold-mask spikes rise where the model reacts to losing the answer glyphs
diff --git a/packages/snapcompact/research/snapcompact_3d_activation_viz.py b/packages/snapcompact/research/snapcompact_3d_activation_viz.py index d002182d2..8fc0608bf 100644 --- a/packages/snapcompact/research/snapcompact_3d_activation_viz.py +++ b/packages/snapcompact/research/snapcompact_3d_activation_viz.py @@ -34,8 +34,12 @@ AMBER = (255, 194, 65) def font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: candidates = [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ] for path in candidates: if path and Path(path).exists(): @@ -80,7 +84,15 @@ def style_3d(ax, title: str, subtitle: str, color: str) -> None: ax.set_box_aspect((2.7, 0.85, 0.55)) -def draw_surface(ax, arr: np.ndarray, cmap_name: str, title: str, subtitle: str, color: str, zmax: float = 1.0) -> None: +def draw_surface( + ax, + arr: np.ndarray, + cmap_name: str, + title: str, + subtitle: str, + color: str, + zmax: float = 1.0, +) -> None: y = np.arange(arr.shape[0]) x = np.arange(arr.shape[1]) X, Y = np.meshgrid(x, y) @@ -99,13 +111,23 @@ def draw_surface(ax, arr: np.ndarray, cmap_name: str, title: str, subtitle: str, alpha=0.98, ) # A dark floor with projected contour lines makes the shape read as 3D. - ax.contour(X, Y, Z, zdir="z", offset=-0.05, levels=9, cmap=cmap, linewidths=0.8, alpha=0.72) + ax.contour( + X, Y, Z, zdir="z", offset=-0.05, levels=9, cmap=cmap, linewidths=0.8, alpha=0.72 + ) ax.set_zlim(-0.05, zmax) ax.set_ylim(arr.shape[0] - 1, 0) style_3d(ax, title, subtitle, color) -def crop_with_box(img: Image.Image, start: int, end: int, cols: int, adv: int, pitch: int, pad_cells: int = 34) -> Image.Image: +def crop_with_box( + img: Image.Image, + start: int, + end: int, + cols: int, + adv: int, + pitch: int, + pad_cells: int = 34, +) -> Image.Image: row0 = max(0, start // cols - 5) row1 = min(img.height // pitch, end // cols + 6) col0 = max(0, start % cols - pad_cells) @@ -122,23 +144,52 @@ def crop_with_box(img: Image.Image, start: int, end: int, cols: int, adv: int, p return crop -def paste_fit(canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int]) -> None: +def paste_fit( + canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int] +) -> None: x0, y0, x1, y1 = box scale = min((x1 - x0) / img.width, (y1 - y0) / img.height) - resized = img.resize((max(1, round(img.width * scale)), max(1, round(img.height * scale))), Image.Resampling.NEAREST) - canvas.paste(resized, (x0 + (x1 - x0 - resized.width) // 2, y0 + (y1 - y0 - resized.height) // 2)) + resized = img.resize( + (max(1, round(img.width * scale)), max(1, round(img.height * scale))), + Image.Resampling.NEAREST, + ) + canvas.paste( + resized, + (x0 + (x1 - x0 - resized.width) // 2, y0 + (y1 - y0 - resized.height) // 2), + ) -def render_matplotlib_panel(answer: np.ndarray, random: np.ndarray, ratio: np.ndarray) -> Image.Image: +def render_matplotlib_panel( + answer: np.ndarray, random: np.ndarray, ratio: np.ndarray +) -> Image.Image: fig = plt.figure(figsize=(16.6, 9.0), dpi=170) fig.patch.set_facecolor("#05070a") - gs = fig.add_gridspec(2, 2, left=0.02, right=0.99, top=0.96, bottom=0.04, wspace=0.03, hspace=0.08) + gs = fig.add_gridspec( + 2, 2, left=0.02, right=0.99, top=0.96, bottom=0.04, wspace=0.03, hspace=0.08 + ) ax1 = fig.add_subplot(gs[0, 0], projection="3d") ax2 = fig.add_subplot(gs[0, 1], projection="3d") ax3 = fig.add_subplot(gs[1, :], projection="3d") - draw_surface(ax1, answer, "magma", "Gold answer mask", "true answer cells erased", "#ff533e") - draw_surface(ax2, random, "viridis", "Random control mask", "same-sized blank elsewhere", "#91ff70") - draw_surface(ax3, ratio, "inferno", "Answer / random ratio", "where the missing answer leaves a larger residual-stream scar", "#ffc241", zmax=1.08) + draw_surface( + ax1, answer, "magma", "Gold answer mask", "true answer cells erased", "#ff533e" + ) + draw_surface( + ax2, + random, + "viridis", + "Random control mask", + "same-sized blank elsewhere", + "#91ff70", + ) + draw_surface( + ax3, + ratio, + "inferno", + "Answer / random ratio", + "where the missing answer leaves a larger residual-stream scar", + "#ffc241", + zmax=1.08, + ) tmp = HERE / "results" / ".snapcompact-3d-panel.png" fig.savefig(tmp, facecolor=fig.get_facecolor(), transparent=False) plt.close(fig) @@ -149,8 +200,12 @@ def render_matplotlib_panel(answer: np.ndarray, random: np.ndarray, ratio: np.nd def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--result-dir", default=str(HERE / "results" / "tensor-heatmap-paddleocr-q7")) - ap.add_argument("--out", default=str(HERE / "results" / "snapcompact-3d-activation-terrain.png")) + ap.add_argument( + "--result-dir", default=str(HERE / "results" / "tensor-heatmap-paddleocr-q7") + ) + ap.add_argument( + "--out", default=str(HERE / "results" / "snapcompact-3d-activation-terrain.png") + ) ap.add_argument("--bins", type=int, default=128) args = ap.parse_args() @@ -171,15 +226,29 @@ def main() -> None: gd = ImageDraw.Draw(glow) gd.ellipse((-260, -160, 950, 600), fill=(255, 83, 62, 34)) gd.ellipse((1100, 220, 2500, 1500), fill=(80, 220, 255, 28)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(82))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(82)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) draw.text((64, 42), "SNAPCOMPACT WHITEBOX", fill=AMBER, font=font(24, True)) - draw.text((64, 82), "Activation terrain from a missing answer", fill=INK, font=font(68, True)) - draw.text((66, 168), "Actual decoder hidden states: layer × image-token bin × ||original − masked||. A blog-friendly 3D tensor slice, not a schematic.", fill=MUTED, font=font(26)) + draw.text( + (64, 82), + "Activation terrain from a missing answer", + fill=INK, + font=font(68, True), + ) + draw.text( + (66, 168), + "Actual decoder hidden states: layer × image-token bin × ||original − masked||. A blog-friendly 3D tensor slice, not a schematic.", + fill=MUTED, + font=font(26), + ) # Left evidence strip. - draw.rounded_rectangle((64, 238, 600, 1234), radius=30, fill=PANEL, outline=(31, 42, 50), width=1) + draw.rounded_rectangle( + (64, 238, 600, 1234), radius=30, fill=PANEL, outline=(31, 42, 50), width=1 + ) q = summary["question"] draw.text((96, 270), "visual intervention", fill=INK, font=font(32, True)) draw.text((96, 310), "answer cells are blanked", fill=MUTED, font=font(18)) @@ -189,10 +258,14 @@ def main() -> None: crop = crop_with_box(base, q["answer_start"], q["answer_end"], cols, 8, 13) masked_crop = crop_with_box(masked, q["answer_start"], q["answer_end"], cols, 8, 13) draw.text((96, 366), "ORIGINAL", fill=CYAN, font=font(17, True)) - draw.rounded_rectangle((96, 394, 568, 560), radius=14, fill=(244, 242, 230), outline=CYAN, width=3) + draw.rounded_rectangle( + (96, 394, 568, 560), radius=14, fill=(244, 242, 230), outline=CYAN, width=3 + ) paste_fit(canvas, crop, (112, 410, 552, 544)) draw.text((96, 618), "ANSWER ERASED", fill=RED, font=font(17, True)) - draw.rounded_rectangle((96, 646, 568, 812), radius=14, fill=(244, 242, 230), outline=RED, width=3) + draw.rounded_rectangle( + (96, 646, 568, 812), radius=14, fill=(244, 242, 230), outline=RED, width=3 + ) paste_fit(canvas, masked_crop, (112, 662, 552, 796)) question = q["q"] if len(question) > 54: @@ -202,15 +275,31 @@ def main() -> None: draw.text((96, 990), "gold answer", fill=MUTED, font=font(16, True)) draw.text((96, 1024), str(q["answer_text"]), fill=AMBER, font=font(42, True)) draw.text((96, 1110), f"{summary['layers']} layers", fill=MUTED, font=font(20)) - draw.text((96, 1142), f"{summary['image_tokens']} image tokens", fill=MUTED, font=font(20)) - draw.text((96, 1174), f"answer/random Δ = {summary['answer_over_random_delta']:.2f}×", fill=INK, font=font(22, True)) + draw.text( + (96, 1142), f"{summary['image_tokens']} image tokens", fill=MUTED, font=font(20) + ) + draw.text( + (96, 1174), + f"answer/random Δ = {summary['answer_over_random_delta']:.2f}×", + fill=INK, + font=font(22, True), + ) # Main 3D panel. - draw.rounded_rectangle((632, 238, 2134, 1234), radius=30, fill=PANEL, outline=(31, 42, 50), width=1) + draw.rounded_rectangle( + (632, 238, 2134, 1234), radius=30, fill=PANEL, outline=(31, 42, 50), width=1 + ) panel = panel.resize((1450, 786), Image.Resampling.LANCZOS) canvas.paste(panel, (660, 330)) - draw.text((672, 268), "3D residual-stream delta terrain", fill=INK, font=font(36, True)) - draw.text((672, 311), "Gold-mask spikes rise where the model’s image-token activations react to losing the answer glyphs.", fill=MUTED, font=font(20)) + draw.text( + (672, 268), "3D residual-stream delta terrain", fill=INK, font=font(36, True) + ) + draw.text( + (672, 311), + "Gold-mask spikes rise where the model’s image-token activations react to losing the answer glyphs.", + fill=MUTED, + font=font(20), + ) # Color scale. cmap = cm.get_cmap("magma") diff --git a/packages/snapcompact/research/snapcompact_activation_probe.py b/packages/snapcompact/research/snapcompact_activation_probe.py index b65036dfe..47d819067 100644 --- a/packages/snapcompact/research/snapcompact_activation_probe.py +++ b/packages/snapcompact/research/snapcompact_activation_probe.py @@ -29,7 +29,11 @@ sys.path.insert(0, str(HERE)) import squad # noqa: E402 from bdf import capacity, render # noqa: E402 from run import CACHE, FONTS, load_prompt # noqa: E402 -from snapcompact_blackbox_occlusion import mask_cells, random_span, sample_answer_questions # noqa: E402 +from snapcompact_blackbox_occlusion import ( + mask_cells, + random_span, + sample_answer_questions, +) # noqa: E402 DEFAULT_MODEL_DIR = ( "/home/can/.cache/huggingface/hub/models--PaddlePaddle--PaddleOCR-VL/" @@ -77,20 +81,36 @@ def to_device(batch: dict[str, Any], device: Any) -> dict[str, Any]: return {k: (v.to(device) if hasattr(v, "to") else v) for k, v in batch.items()} -def hidden_features(model: Any, processor: Any, *, image: Image.Image | None, text: str, device: Any) -> list[np.ndarray]: +def hidden_features( + model: Any, processor: Any, *, image: Image.Image | None, text: str, device: Any +) -> list[np.ndarray]: import torch if image is None: messages = [{"role": "user", "content": [{"type": "text", "text": text}]}] - templated = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) + templated = processor.apply_chat_template( + messages, tokenize=False, add_generation_prompt=True + ) batch = processor(text=templated, return_tensors="pt") else: - messages = [{"role": "user", "content": [{"type": "image", "image": image}, {"type": "text", "text": text}]}] - templated = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) + messages = [ + { + "role": "user", + "content": [ + {"type": "image", "image": image}, + {"type": "text", "text": text}, + ], + } + ] + templated = processor.apply_chat_template( + messages, tokenize=False, add_generation_prompt=True + ) batch = processor(images=image, text=templated, return_tensors="pt") batch = to_device(batch, device) with torch.no_grad(): - out = model(**batch, output_hidden_states=True, output_attentions=False, use_cache=False) + out = model( + **batch, output_hidden_states=True, output_attentions=False, use_cache=False + ) feats: list[np.ndarray] = [] for h in out.hidden_states: # Mean-pool the prompt sequence. This avoids brittle alignment between @@ -136,19 +156,38 @@ def main() -> None: fill = (255, 255, 255) if args.variant not in ("dark", "dark-sent") else (0, 0, 0) print(f"loading {args.model_dir}", flush=True) - processor = AutoProcessor.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False) + processor = AutoProcessor.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False + ) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") dtype = torch.bfloat16 if device.type == "cuda" else torch.float32 - model = AutoModel.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, dtype=dtype).to(device).eval() + model = ( + AutoModel.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True, dtype=dtype + ) + .to(device) + .eval() + ) - feature_sets: dict[str, list[list[np.ndarray]]] = {"text": [], "image": [], "answer_mask": [], "random_mask": []} + feature_sets: dict[str, list[list[np.ndarray]]] = { + "text": [], + "image": [], + "answer_mask": [], + "random_mask": [], + } records: list[dict[str, Any]] = [] for qi, q in enumerate(questions): span_len = max(1, q["answer_end"] - q["answer_start"]) rng = random.Random(args.seed * 31 + qi) - rand_start, rand_end = random_span(rng, len(chunk), span_len, q["answer_start"], q["answer_end"]) - answer_img = mask_cells(base_img, q["answer_start"], q["answer_end"], cols, cfg.adv, cfg.pitch, fill) - random_img = mask_cells(base_img, rand_start, rand_end, cols, cfg.adv, cfg.pitch, fill) + rand_start, rand_end = random_span( + rng, len(chunk), span_len, q["answer_start"], q["answer_end"] + ) + answer_img = mask_cells( + base_img, q["answer_start"], q["answer_end"], cols, cfg.adv, cfg.pitch, fill + ) + random_img = mask_cells( + base_img, rand_start, rand_end, cols, cfg.adv, cfg.pitch, fill + ) answer_path = img_dir / f"q{qi}-answer-mask.png" random_path = img_dir / f"q{qi}-random-mask.png" answer_img.save(answer_path) @@ -160,10 +199,26 @@ def main() -> None: f"{chunk}\n\nQuestion: {q['q']}\n" "Answer with only the shortest extractive answer." ) - feature_sets["text"].append(hidden_features(model, processor, image=None, text=text_prompt, device=device)) - feature_sets["image"].append(hidden_features(model, processor, image=base_img, text=img_prompt, device=device)) - feature_sets["answer_mask"].append(hidden_features(model, processor, image=answer_img, text=img_prompt, device=device)) - feature_sets["random_mask"].append(hidden_features(model, processor, image=random_img, text=img_prompt, device=device)) + feature_sets["text"].append( + hidden_features( + model, processor, image=None, text=text_prompt, device=device + ) + ) + feature_sets["image"].append( + hidden_features( + model, processor, image=base_img, text=img_prompt, device=device + ) + ) + feature_sets["answer_mask"].append( + hidden_features( + model, processor, image=answer_img, text=img_prompt, device=device + ) + ) + feature_sets["random_mask"].append( + hidden_features( + model, processor, image=random_img, text=img_prompt, device=device + ) + ) records.append( { "question_index": qi, @@ -201,7 +256,11 @@ def main() -> None: "cos_image_random_mask": paired_cosine(img, rnd), "answer_delta_norm": float(answer_delta.mean()), "random_delta_norm": float(random_delta.mean()), - "answer_over_random_delta": float(answer_delta.mean() / random_delta.mean()) if random_delta.mean() else float("inf"), + "answer_over_random_delta": float( + answer_delta.mean() / random_delta.mean() + ) + if random_delta.mean() + else float("inf"), } ) diff --git a/packages/snapcompact/research/snapcompact_blackbox_occlusion.py b/packages/snapcompact/research/snapcompact_blackbox_occlusion.py index 15faac599..a1077b876 100644 --- a/packages/snapcompact/research/snapcompact_blackbox_occlusion.py +++ b/packages/snapcompact/research/snapcompact_blackbox_occlusion.py @@ -33,13 +33,17 @@ from bdf import capacity, render # noqa: E402 from run import CACHE, FONTS, load_prompt, sha8 # noqa: E402 -def sample_answer_questions(paras: list[dict], offsets: list[int], start: int, end: int, n: int, seed: int) -> list[dict]: +def sample_answer_questions( + paras: list[dict], offsets: list[int], start: int, end: int, n: int, seed: int +) -> list[dict]: """Sample questions like squad.sample_chunk_questions, preserving answer offsets.""" rng = random.Random(seed * 1_000_003 + start) eligible = [ i for i in range(len(offsets)) - if offsets[i] >= start and offsets[i] + len(paras[i]["ctx"]) <= end and paras[i].get("qas") + if offsets[i] >= start + and offsets[i] + len(paras[i]["ctx"]) <= end + and paras[i].get("qas") ] if not eligible: return [] @@ -59,14 +63,25 @@ def sample_answer_questions(paras: list[dict], offsets: list[int], start: int, e "golds": sorted({a["text"] for a in answers}), "answer_text": answer["text"], "answer_start": offsets[pi] - start + int(answer["answer_start"]), - "answer_end": offsets[pi] - start + int(answer["answer_start"]) + len(answer["text"]), + "answer_end": offsets[pi] + - start + + int(answer["answer_start"]) + + len(answer["text"]), "pos_rel": (offsets[pi] - start) / (end - start), } ) return picked -def mask_cells(img: Image.Image, start: int, end: int, cols: int, adv: int, pitch: int, fill: tuple[int, int, int]) -> Image.Image: +def mask_cells( + img: Image.Image, + start: int, + end: int, + cols: int, + adv: int, + pitch: int, + fill: tuple[int, int, int], +) -> Image.Image: out = img.copy() draw = ImageDraw.Draw(out) start = max(0, start) @@ -84,7 +99,9 @@ def mask_cells(img: Image.Image, start: int, end: int, cols: int, adv: int, pitc return out -def random_span(rng: random.Random, text_len: int, span_len: int, avoid_start: int, avoid_end: int) -> tuple[int, int]: +def random_span( + rng: random.Random, text_len: int, span_len: int, avoid_start: int, avoid_end: int +) -> tuple[int, int]: if text_len <= span_len: return 0, text_len for _ in range(100): @@ -96,7 +113,15 @@ def random_span(rng: random.Random, text_len: int, span_len: int, avoid_start: i return start, min(text_len, start + span_len) -def post_chat(endpoint: str, model: str, image_path: Path, prompt: str, max_tokens: int, cache_dir: Path, fresh: bool) -> tuple[str, dict]: +def post_chat( + endpoint: str, + model: str, + image_path: Path, + prompt: str, + max_tokens: int, + cache_dir: Path, + fresh: bool, +) -> tuple[str, dict]: payload_key = sha8(model, prompt, hashlib.sha1(image_path.read_bytes()).hexdigest()) cache_path = cache_dir / f"{payload_key}.json" if cache_path.exists() and not fresh: @@ -111,7 +136,10 @@ def post_chat(endpoint: str, model: str, image_path: Path, prompt: str, max_toke "role": "user", "content": [ {"type": "text", "text": prompt}, - {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{image_b64}"}}, + { + "type": "image_url", + "image_url": {"url": f"data:image/png;base64,{image_b64}"}, + }, ], } ], @@ -147,19 +175,27 @@ def aggregate(records: list[dict]) -> dict: "em": sum(r["em"] for r in rows) / max(1, len(rows)), "f1": sum(r["f1"] for r in rows) / max(1, len(rows)), "abstained": sum(1 for r in rows if "unreadable" in r["answer"].lower()), - "prompt_tokens": sum((r.get("usage") or {}).get("prompt_tokens", 0) for r in rows), - "completion_tokens": sum((r.get("usage") or {}).get("completion_tokens", 0) for r in rows), + "prompt_tokens": sum( + (r.get("usage") or {}).get("prompt_tokens", 0) for r in rows + ), + "completion_tokens": sum( + (r.get("usage") or {}).get("completion_tokens", 0) for r in rows + ), } base = out["variants"].get("original", {}).get("f1", 0.0) out["drops"] = { - name: base - row["f1"] for name, row in out["variants"].items() if name != "original" + name: base - row["f1"] + for name, row in out["variants"].items() + if name != "original" } return out def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--endpoint", default="http://spark.internal:8000/v1/chat/completions") + ap.add_argument( + "--endpoint", default="http://spark.internal:8000/v1/chat/completions" + ) ap.add_argument("--model", default="Qwen2.5-VL-7B-Instruct-NVFP4") ap.add_argument("--font", default="5x8", choices=sorted(FONTS)) ap.add_argument("--variant", default="bw") @@ -184,13 +220,17 @@ def main() -> None: paras = squad.load_paragraphs(CACHE)[: args.limit_paras] flow, offsets = squad.build_flow(paras) prompt_base = load_prompt("qa-image.md").format(cols=cols, rows=rows) - mask_fill = (255, 255, 255) if args.variant not in ("dark", "dark-sent") else (0, 0, 0) + mask_fill = ( + (255, 255, 255) if args.variant not in ("dark", "dark-sent") else (0, 0, 0) + ) tasks = [] for start in range(0, len(flow), budget): end = min(start + budget, len(flow)) chunk = flow[start:end] - questions = sample_answer_questions(paras, offsets, start, end, args.qpc, args.seed) + questions = sample_answer_questions( + paras, offsets, start, end, args.qpc, args.seed + ) if questions: tasks.append((start, end, chunk, questions)) @@ -203,13 +243,25 @@ def main() -> None: for qi, q in enumerate(questions): span_len = max(1, q["answer_end"] - q["answer_start"]) rng = random.Random(args.seed * 17 + start + qi) - rand_start, rand_end = random_span(rng, len(chunk), span_len, q["answer_start"], q["answer_end"]) + rand_start, rand_end = random_span( + rng, len(chunk), span_len, q["answer_start"], q["answer_end"] + ) answer_path = img_dir / f"chunk-{start}-q{qi}-answer-mask.png" random_path = img_dir / f"chunk-{start}-q{qi}-random-mask.png" if not answer_path.exists(): - mask_cells(base_img, q["answer_start"], q["answer_end"], cols, cfg.adv, cfg.pitch, mask_fill).save(answer_path) + mask_cells( + base_img, + q["answer_start"], + q["answer_end"], + cols, + cfg.adv, + cfg.pitch, + mask_fill, + ).save(answer_path) if not random_path.exists(): - mask_cells(base_img, rand_start, rand_end, cols, cfg.adv, cfg.pitch, mask_fill).save(random_path) + mask_cells( + base_img, rand_start, rand_end, cols, cfg.adv, cfg.pitch, mask_fill + ).save(random_path) prompt = ( f"{prompt_base}\n\nQuestion: {q['q']}\n" @@ -221,7 +273,15 @@ def main() -> None: ("answer_mask", answer_path), ("random_mask", random_path), ): - answer, usage = post_chat(args.endpoint, args.model, path, prompt, args.max_tokens, cache_dir, args.fresh) + answer, usage = post_chat( + args.endpoint, + args.model, + path, + prompt, + args.max_tokens, + cache_dir, + args.fresh, + ) records.append( { "chunk": start, @@ -242,7 +302,10 @@ def main() -> None: "usage": usage, } ) - print(f"{len(records):04d} {variant_name:<11} f1={records[-1]['f1']:.3f} answer={answer[:80]!r}", flush=True) + print( + f"{len(records):04d} {variant_name:<11} f1={records[-1]['f1']:.3f} answer={answer[:80]!r}", + flush=True, + ) summary = { "args": vars(args), diff --git a/packages/snapcompact/research/snapcompact_blog_viz.py b/packages/snapcompact/research/snapcompact_blog_viz.py index c431c2205..ba29acc58 100644 --- a/packages/snapcompact/research/snapcompact_blog_viz.py +++ b/packages/snapcompact/research/snapcompact_blog_viz.py @@ -32,9 +32,13 @@ PALETTE = { def font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: candidates = [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", "/System/Library/Fonts/Monaco.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ] for path in candidates: if path and Path(path).exists(): @@ -42,11 +46,25 @@ def font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.Im return ImageFont.load_default() -def rounded(draw: ImageDraw.ImageDraw, xy: tuple[int, int, int, int], fill: tuple[int, int, int], outline=None, radius=24, width=1) -> None: +def rounded( + draw: ImageDraw.ImageDraw, + xy: tuple[int, int, int, int], + fill: tuple[int, int, int], + outline=None, + radius=24, + width=1, +) -> None: draw.rounded_rectangle(xy, radius=radius, fill=fill, outline=outline, width=width) -def draw_label(draw: ImageDraw.ImageDraw, xy: tuple[int, int], text: str, color: tuple[int, int, int], size: int = 24, bold: bool = False) -> None: +def draw_label( + draw: ImageDraw.ImageDraw, + xy: tuple[int, int], + text: str, + color: tuple[int, int, int], + size: int = 24, + bold: bool = False, +) -> None: draw.text(xy, text, fill=color, font=font(size, bold=bold)) @@ -68,7 +86,9 @@ def chart( y = gy0 + round((gy1 - gy0) * i / 4) draw.line((gx0, y, gx1, y), fill=PALETTE["grid"], width=1) value = y_max - (y_max - y_min) * i / 4 - draw.text((x0 + 18, y - 9), f"{value:.2f}", fill=PALETTE["muted"], font=font(13)) + draw.text( + (x0 + 18, y - 9), f"{value:.2f}", fill=PALETTE["muted"], font=font(13) + ) n = len(series[0][1]) for label, values, color in series: pts = [] @@ -87,7 +107,15 @@ def chart( lx += 210 -def crop_with_box(img: Image.Image, start: int, end: int, cols: int, adv: int, pitch: int, pad_cells: int = 24) -> Image.Image: +def crop_with_box( + img: Image.Image, + start: int, + end: int, + cols: int, + adv: int, + pitch: int, + pad_cells: int = 24, +) -> Image.Image: row0 = max(0, start // cols - 4) row1 = min(img.height // pitch, end // cols + 5) col0 = max(0, start % cols - pad_cells) @@ -103,14 +131,21 @@ def crop_with_box(img: Image.Image, start: int, end: int, cols: int, adv: int, p bx1 = min(crop.width - 1, ((end - 1) % cols - col0 + 2) * adv) by0 = max(0, (start // cols - row0) * pitch - 1) by1 = min(crop.height - 1, ((end - 1) // cols - row0 + 1) * pitch + 1) - draw.rounded_rectangle((bx0, by0, bx1, by1), radius=3, outline=PALETTE["red"], width=3) + draw.rounded_rectangle( + (bx0, by0, bx1, by1), radius=3, outline=PALETTE["red"], width=3 + ) return crop -def paste_fit(canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int]) -> None: +def paste_fit( + canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int] +) -> None: x0, y0, x1, y1 = box scale = min((x1 - x0) / img.width, (y1 - y0) / img.height) - resized = img.resize((max(1, round(img.width * scale)), max(1, round(img.height * scale))), Image.Resampling.NEAREST) + resized = img.resize( + (max(1, round(img.width * scale)), max(1, round(img.height * scale))), + Image.Resampling.NEAREST, + ) px = x0 + (x1 - x0 - resized.width) // 2 py = y0 + (y1 - y0 - resized.height) // 2 canvas.paste(resized, (px, py)) @@ -118,16 +153,26 @@ def paste_fit(canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, i def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--activation", default=str(HERE / "results" / "activation-paddleocr-8x13-n16")) - ap.add_argument("--occlusion", default=str(HERE / "results" / "snapcompact-occlusion-qwen-8x13")) - ap.add_argument("--out", default=str(HERE / "results" / "snapcompact-blog-whitebox.png")) + ap.add_argument( + "--activation", default=str(HERE / "results" / "activation-paddleocr-8x13-n16") + ) + ap.add_argument( + "--occlusion", default=str(HERE / "results" / "snapcompact-occlusion-qwen-8x13") + ) + ap.add_argument( + "--out", default=str(HERE / "results" / "snapcompact-blog-whitebox.png") + ) args = ap.parse_args() act_dir = Path(args.activation) occ_dir = Path(args.occlusion) summary = json.loads((act_dir / "summary.json").read_text()) occ = json.loads((occ_dir / "summary.json").read_text()) - records = [json.loads(line) for line in (act_dir / "records.jsonl").read_text().splitlines() if line] + records = [ + json.loads(line) + for line in (act_dir / "records.jsonl").read_text().splitlines() + if line + ] layers = summary["layers"] w, h = 1800, 1040 @@ -141,23 +186,58 @@ def main() -> None: gd = ImageDraw.Draw(glow) gd.ellipse((-160, -220, 760, 520), fill=(255, 104, 72, 34)) gd.ellipse((1100, 130, 2100, 1160), fill=(67, 210, 255, 28)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(70))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(70)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) - draw_label(draw, (56, 38), "SNAPCOMPACT UNDER THE MICROSCOPE", PALETTE["amber"], 21, True) - draw_label(draw, (56, 78), "Dense text-images leave a white-box trace", PALETTE["text"], 54, True) - draw_label(draw, (58, 145), "Text and image prompts converge late; blanking the gold answer region perturbs hidden states far more than an equal random blank.", PALETTE["muted"], 24) + draw_label( + draw, (56, 38), "SNAPCOMPACT UNDER THE MICROSCOPE", PALETTE["amber"], 21, True + ) + draw_label( + draw, + (56, 78), + "Dense text-images leave a white-box trace", + PALETTE["text"], + 54, + True, + ) + draw_label( + draw, + (58, 145), + "Text and image prompts converge late; blanking the gold answer region perturbs hidden states far more than an equal random blank.", + PALETTE["muted"], + 24, + ) # Big stat cards. stats = [ - ("Qwen black-box F1", f"{occ['variants']['original']['f1']:.2f}", "original image"), - ("gold-mask drop", f"−{occ['drops']['answer_mask']:.2f}", "answer region blanked"), - ("random-mask drop", f"−{occ['drops']['random_mask']:.2f}", "same-size random blank"), + ( + "Qwen black-box F1", + f"{occ['variants']['original']['f1']:.2f}", + "original image", + ), + ( + "gold-mask drop", + f"−{occ['drops']['answer_mask']:.2f}", + "answer region blanked", + ), + ( + "random-mask drop", + f"−{occ['drops']['random_mask']:.2f}", + "same-size random blank", + ), ] sx = 56 card_w = 258 for title, value, caption in stats: - rounded(draw, (sx, 205, sx + card_w, 330), PALETTE["panel2"], outline=(36, 48, 56), radius=22) + rounded( + draw, + (sx, 205, sx + card_w, 330), + PALETTE["panel2"], + outline=(36, 48, 56), + radius=22, + ) draw_label(draw, (sx + 20, 226), title, PALETTE["muted"], 16) draw_label(draw, (sx + 20, 252), value, PALETTE["text"], 42, True) draw_label(draw, (sx + 20, 300), caption, PALETTE["muted"], 14) @@ -167,9 +247,21 @@ def main() -> None: draw, (56, 368, 872, 668), [ - ("text ↔ image CKA", [x["cka_text_image"] for x in layers], PALETTE["accent2"]), - ("answer-mask CKA", [x["cka_image_answer_mask"] for x in layers], PALETTE["accent"]), - ("random-mask CKA", [x["cka_image_random_mask"] for x in layers], PALETTE["green"]), + ( + "text ↔ image CKA", + [x["cka_text_image"] for x in layers], + PALETTE["accent2"], + ), + ( + "answer-mask CKA", + [x["cka_image_answer_mask"] for x in layers], + PALETTE["accent"], + ), + ( + "random-mask CKA", + [x["cka_image_random_mask"] for x in layers], + PALETTE["green"], + ), ], 0.2, 1.0, @@ -180,7 +272,11 @@ def main() -> None: draw, (56, 698, 872, 990), [ - ("answer / random perturbation", [x["answer_over_random_delta"] for x in layers], PALETTE["amber"]), + ( + "answer / random perturbation", + [x["answer_over_random_delta"] for x in layers], + PALETTE["amber"], + ), ], 1.0, 1.6, @@ -191,26 +287,63 @@ def main() -> None: # Visual crop panel. panel = (920, 205, 1744, 990) rounded(draw, panel, PALETTE["panel"], outline=(35, 47, 56), radius=26) - draw_label(draw, (950, 232), "What the mask test looks like", PALETTE["text"], 34, True) - draw_label(draw, (950, 274), "Same question, same bitmap. Only the gold answer cells are erased.", PALETTE["muted"], 19) + draw_label( + draw, (950, 232), "What the mask test looks like", PALETTE["text"], 34, True + ) + draw_label( + draw, + (950, 274), + "Same question, same bitmap. Only the gold answer cells are erased.", + PALETTE["muted"], + 19, + ) base = Image.open(act_dir / "activation-images" / "base.png").convert("RGB") # Use a later question if possible because it gives a better-looking crop. rec = records[min(7, len(records) - 1)] - ans = Image.open(act_dir / "activation-images" / f"q{rec['question_index']}-answer-mask.png").convert("RGB") - rnd = Image.open(act_dir / "activation-images" / f"q{rec['question_index']}-random-mask.png").convert("RGB") + ans = Image.open( + act_dir / "activation-images" / f"q{rec['question_index']}-answer-mask.png" + ).convert("RGB") + rnd = Image.open( + act_dir / "activation-images" / f"q{rec['question_index']}-random-mask.png" + ).convert("RGB") cols = summary["geometry"]["cols"] adv = 8 pitch = 13 crops = [ - ("original", crop_with_box(base, rec["answer_start"], rec["answer_end"], cols, adv, pitch), PALETTE["accent2"]), - ("answer masked", crop_with_box(ans, rec["answer_start"], rec["answer_end"], cols, adv, pitch), PALETTE["accent"]), - ("random masked", crop_with_box(rnd, rec["random_start"], rec["random_end"], cols, adv, pitch), PALETTE["green"]), + ( + "original", + crop_with_box( + base, rec["answer_start"], rec["answer_end"], cols, adv, pitch + ), + PALETTE["accent2"], + ), + ( + "answer masked", + crop_with_box( + ans, rec["answer_start"], rec["answer_end"], cols, adv, pitch + ), + PALETTE["accent"], + ), + ( + "random masked", + crop_with_box( + rnd, rec["random_start"], rec["random_end"], cols, adv, pitch + ), + PALETTE["green"], + ), ] y = 334 for label, img, color in crops: draw_label(draw, (950, y - 31), label.upper(), color, 17, True) - rounded(draw, (950, y, 1714, y + 150), (244, 242, 230), outline=color, radius=14, width=3) + rounded( + draw, + (950, y, 1714, y + 150), + (244, 242, 230), + outline=color, + radius=14, + width=3, + ) paste_fit(canvas, img, (966, y + 16, 1698, y + 134)) y += 198 @@ -219,7 +352,14 @@ def main() -> None: q = q[:89] + "…" draw_label(draw, (950, 916), "sample question", PALETTE["muted"], 17, True) draw_label(draw, (950, 940), q, PALETTE["text"], 20) - draw_label(draw, (950, 966), f"gold answer: {rec['answer_text']}", PALETTE["amber"], 18, True) + draw_label( + draw, + (950, 966), + f"gold answer: {rec['answer_text']}", + PALETTE["amber"], + 18, + True, + ) out = Path(args.out) out.parent.mkdir(parents=True, exist_ok=True) diff --git a/packages/snapcompact/research/snapcompact_carrier_convergence.py b/packages/snapcompact/research/snapcompact_carrier_convergence.py index cfdb21af2..d3be75a7b 100644 --- a/packages/snapcompact/research/snapcompact_carrier_convergence.py +++ b/packages/snapcompact/research/snapcompact_carrier_convergence.py @@ -44,10 +44,15 @@ def make_text_prompt(chunk: str, question: str) -> str: def make_image_prompt(cols: int, rows: int, question: str) -> str: - return load_prompt("qa-image.md").format(cols=cols, rows=rows) + f"\n\nQuestion: {question}\nAnswer with only the shortest extractive answer." + return ( + load_prompt("qa-image.md").format(cols=cols, rows=rows) + + f"\n\nQuestion: {question}\nAnswer with only the shortest extractive answer." + ) -def capture_last_token(model: Any, processor: Any, device: Any, text: str, image: Image.Image | None) -> tuple[np.ndarray, str]: +def capture_last_token( + model: Any, processor: Any, device: Any, text: str, image: Image.Image | None +) -> tuple[np.ndarray, str]: """Return per-layer hidden state at the final prompt position plus a short generation.""" import torch @@ -55,7 +60,11 @@ def capture_last_token(model: Any, processor: Any, device: Any, text: str, image if image is not None: content.append({"type": "image", "image": image}) content.append({"type": "text", "text": text}) - templated = processor.apply_chat_template([{"role": "user", "content": content}], tokenize=False, add_generation_prompt=True) + templated = processor.apply_chat_template( + [{"role": "user", "content": content}], + tokenize=False, + add_generation_prompt=True, + ) if image is not None: batch = processor(images=image, text=templated, return_tensors="pt") else: @@ -64,8 +73,12 @@ def capture_last_token(model: Any, processor: Any, device: Any, text: str, image with torch.no_grad(): out = model(**batch, output_hidden_states=True, use_cache=False) generated = model.generate(**batch, max_new_tokens=16, do_sample=False) - states = np.stack([h[0, -1, :].float().detach().cpu().numpy() for h in out.hidden_states], axis=0) - answer = processor.batch_decode(generated[:, batch["input_ids"].shape[1] :], skip_special_tokens=True)[0].strip() + states = np.stack( + [h[0, -1, :].float().detach().cpu().numpy() for h in out.hidden_states], axis=0 + ) + answer = processor.batch_decode( + generated[:, batch["input_ids"].shape[1] :], skip_special_tokens=True + )[0].strip() return states.astype(np.float32, copy=False), answer @@ -100,7 +113,9 @@ def main() -> None: paras = squad.load_paragraphs(CACHE)[: args.limit_paras] flow, offsets = squad.build_flow(paras) chunk = flow[: min(len(flow), budget)] - questions = sample_answer_questions(paras, offsets, 0, len(chunk), args.questions * 2, args.seed) + questions = sample_answer_questions( + paras, offsets, 0, len(chunk), args.questions * 2, args.seed + ) # Deduplicate gold answers so the RSA geometry has distinct content per row. seen: set[str] = set() picked: list[dict[str, Any]] = [] @@ -117,16 +132,28 @@ def main() -> None: img.save(img_dir / "image-carrier.png") print(f"loading {args.model_dir}", flush=True) - processor = AutoProcessor.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False) - model = Qwen2_5_VLForConditionalGeneration.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, dtype=torch.bfloat16, device_map="auto").eval() + processor = AutoProcessor.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False + ) + model = Qwen2_5_VLForConditionalGeneration.from_pretrained( + args.model_dir, + local_files_only=True, + trust_remote_code=True, + dtype=torch.bfloat16, + device_map="auto", + ).eval() device = next(model.parameters()).device text_states: list[np.ndarray] = [] image_states: list[np.ndarray] = [] records: list[dict[str, Any]] = [] for qi, q in enumerate(picked): - t_states, t_answer = capture_last_token(model, processor, device, make_text_prompt(chunk, q["q"]), None) - i_states, i_answer = capture_last_token(model, processor, device, make_image_prompt(cols, rows, q["q"]), img) + t_states, t_answer = capture_last_token( + model, processor, device, make_text_prompt(chunk, q["q"]), None + ) + i_states, i_answer = capture_last_token( + model, processor, device, make_image_prompt(cols, rows, q["q"]), img + ) text_states.append(t_states) image_states.append(i_states) records.append( @@ -142,7 +169,10 @@ def main() -> None: "agree": squad.f1(t_answer, [i_answer]) >= 0.99, } ) - print(f"{qi + 1}/{len(picked)} text={t_answer!r} image={i_answer!r} gold={q['answer_text']!r}", flush=True) + print( + f"{qi + 1}/{len(picked)} text={t_answer!r} image={i_answer!r} gold={q['answer_text']!r}", + flush=True, + ) text_arr = np.stack(text_states, axis=0) # [Q, L, D] image_arr = np.stack(image_states, axis=0) @@ -173,7 +203,9 @@ def main() -> None: "mismatched_cosine": mismatched, "separation": matched - mismatched, "rsa_pearson": rsa, - "match_rank_accuracy": float((np.argmax(cross, axis=1) == np.arange(n_q)).mean()), + "match_rank_accuracy": float( + (np.argmax(cross, axis=1) == np.arange(n_q)).mean() + ), } ) text_sim_by_layer[layer] = text_sim @@ -204,7 +236,12 @@ def main() -> None: cross_sim=cross_sim_by_layer, ) (out_dir / "summary.json").write_text(json.dumps(summary, indent=1)) - print(json.dumps({k: v for k, v in summary.items() if k not in ("per_layer", "records")}, indent=1)) + print( + json.dumps( + {k: v for k, v in summary.items() if k not in ("per_layer", "records")}, + indent=1, + ) + ) print(f"results -> {out_dir}") diff --git a/packages/snapcompact/research/snapcompact_convergence_3d.py b/packages/snapcompact/research/snapcompact_convergence_3d.py index ec4f8ed28..2c5678e6d 100644 --- a/packages/snapcompact/research/snapcompact_convergence_3d.py +++ b/packages/snapcompact/research/snapcompact_convergence_3d.py @@ -12,6 +12,7 @@ import math from pathlib import Path import matplotlib + matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np @@ -29,8 +30,12 @@ ORANGE = (255, 112, 72) def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: for path in [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ]: if Path(path).exists(): return ImageFont.truetype(path, size) @@ -57,9 +62,14 @@ def smooth_path(path: np.ndarray, passes: int = 2) -> np.ndarray: return out -def render_strands(text_arr: np.ndarray, image_arr: np.ndarray, best_layer: int) -> Image.Image: +def render_strands( + text_arr: np.ndarray, image_arr: np.ndarray, best_layer: int +) -> Image.Image: n_q, n_layers, _ = text_arr.shape - ref = np.concatenate([center(text_arr[:, best_layer, :]), center(image_arr[:, best_layer, :])], axis=0) + ref = np.concatenate( + [center(text_arr[:, best_layer, :]), center(image_arr[:, best_layer, :])], + axis=0, + ) _, _, vt = np.linalg.svd(ref, full_matrices=False) basis = vt[:2].T @@ -86,10 +96,22 @@ def render_strands(text_arr: np.ndarray, image_arr: np.ndarray, best_layer: int) layers_axis = np.arange(n_layers) for qi in range(n_q): color = question_hue(qi, n_q) - tp = smooth_path(np.column_stack([layers_axis, t_proj[qi, :, 0], t_proj[qi, :, 1]])) - ip = smooth_path(np.column_stack([layers_axis, i_proj[qi, :, 0], i_proj[qi, :, 1]])) + tp = smooth_path( + np.column_stack([layers_axis, t_proj[qi, :, 0], t_proj[qi, :, 1]]) + ) + ip = smooth_path( + np.column_stack([layers_axis, i_proj[qi, :, 0], i_proj[qi, :, 1]]) + ) ax.plot(tp[:, 0], tp[:, 1], tp[:, 2], color=color, linewidth=2.6, alpha=0.95) - ax.plot(ip[:, 0], ip[:, 1], ip[:, 2], color=color, linewidth=2.6, alpha=0.55, linestyle=(0, (4, 2))) + ax.plot( + ip[:, 0], + ip[:, 1], + ip[:, 2], + color=color, + linewidth=2.6, + alpha=0.55, + linestyle=(0, (4, 2)), + ) # tie-lines every few layers showing the closing gap for layer in range(1, n_layers, 4): ax.plot( @@ -100,13 +122,41 @@ def render_strands(text_arr: np.ndarray, image_arr: np.ndarray, best_layer: int) linewidth=0.9, alpha=0.38, ) - ax.scatter([0], [t_proj[qi, 0, 0]], [t_proj[qi, 0, 1]], color=color, s=26, marker="o", depthshade=False) - ax.scatter([0], [i_proj[qi, 0, 0]], [i_proj[qi, 0, 1]], color=color, s=30, marker="D", depthshade=False) - ax.scatter([best_layer], [t_proj[qi, best_layer, 0]], [t_proj[qi, best_layer, 1]], color=color, s=46, marker="o", edgecolors="white", linewidths=0.6, depthshade=False) + ax.scatter( + [0], + [t_proj[qi, 0, 0]], + [t_proj[qi, 0, 1]], + color=color, + s=26, + marker="o", + depthshade=False, + ) + ax.scatter( + [0], + [i_proj[qi, 0, 0]], + [i_proj[qi, 0, 1]], + color=color, + s=30, + marker="D", + depthshade=False, + ) + ax.scatter( + [best_layer], + [t_proj[qi, best_layer, 0]], + [t_proj[qi, best_layer, 1]], + color=color, + s=46, + marker="o", + edgecolors="white", + linewidths=0.6, + depthshade=False, + ) # Peak-layer plane. yy, zz = np.meshgrid(np.linspace(-1.05, 1.05, 2), np.linspace(-1.05, 1.05, 2)) - ax.plot_surface(np.full_like(yy, best_layer), yy, zz, color=(1.0, 0.77, 0.27, 0.10), shade=False) + ax.plot_surface( + np.full_like(yy, best_layer), yy, zz, color=(1.0, 0.77, 0.27, 0.10), shade=False + ) ax.set_xlim(0, n_layers - 1) ax.set_ylim(-1.1, 1.1) @@ -118,7 +168,13 @@ def render_strands(text_arr: np.ndarray, image_arr: np.ndarray, best_layer: int) ax.set_box_aspect((2.9, 1.0, 0.9)) tmp = HERE / "results" / ".convergence-3d-panel.png" fig.subplots_adjust(left=0, right=1, top=1, bottom=0) - fig.savefig(tmp, facecolor=fig.get_facecolor(), transparent=False, bbox_inches="tight", pad_inches=0.05) + fig.savefig( + tmp, + facecolor=fig.get_facecolor(), + transparent=False, + bbox_inches="tight", + pad_inches=0.05, + ) plt.close(fig) img = Image.open(tmp).convert("RGB") tmp.unlink(missing_ok=True) @@ -127,8 +183,18 @@ def render_strands(text_arr: np.ndarray, image_arr: np.ndarray, best_layer: int) def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--result-dir", default=str(HERE / "results" / "qwen-carrier-convergence-n12")) - ap.add_argument("--out", default=str(HERE / "results" / "qwen-carrier-convergence-n12" / "convergence-strands-3d.png")) + ap.add_argument( + "--result-dir", default=str(HERE / "results" / "qwen-carrier-convergence-n12") + ) + ap.add_argument( + "--out", + default=str( + HERE + / "results" + / "qwen-carrier-convergence-n12" + / "convergence-strands-3d.png" + ), + ) args = ap.parse_args() result_dir = Path(args.result_dir) summary = json.loads((result_dir / "summary.json").read_text()) @@ -146,11 +212,23 @@ def main() -> None: gd = ImageDraw.Draw(glow) gd.ellipse((-260, -220, 900, 700), fill=(75, 220, 255, 27)) gd.ellipse((1240, 160, 2460, 1360), fill=(255, 112, 72, 25)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(84))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(84)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) - draw.text((64, 42), "QWEN CARRIER CONVERGENCE — 3D STRANDS", fill=AMBER, font=ui_font(24, True)) - draw.text((64, 84), "Twelve thoughts, two doors, one room", fill=INK, font=ui_font(64, True)) + draw.text( + (64, 42), + "QWEN CARRIER CONVERGENCE — 3D STRANDS", + fill=AMBER, + font=ui_font(24, True), + ) + draw.text( + (64, 84), + "Twelve thoughts, two doors, one room", + fill=INK, + font=ui_font(64, True), + ) draw.text( (66, 164), "Each color is one question travelling through the decoder. Solid strand entered as text; dashed strand entered as pixels. Strand pairs braid together by depth.", @@ -158,10 +236,17 @@ def main() -> None: font=ui_font(23), ) - draw.rounded_rectangle((64, 234, 2136, 1146), radius=30, fill=PANEL, outline=(35, 49, 59), width=1) + draw.rounded_rectangle( + (64, 234, 2136, 1146), radius=30, fill=PANEL, outline=(35, 49, 59), width=1 + ) panel = panel.resize((1980, 832), Image.Resampling.LANCZOS) canvas.paste(panel, (104, 286)) - draw.text((96, 252), f"PCA frame fixed at peak layer {best_layer}; per-layer scale normalized", fill=MUTED, font=ui_font(17)) + draw.text( + (96, 252), + f"PCA frame fixed at peak layer {best_layer}; per-layer scale normalized", + fill=MUTED, + font=ui_font(17), + ) stats = [ ("matched cosine", f"{best['matched_cosine']:.2f}"), @@ -171,11 +256,22 @@ def main() -> None: ] sx = 64 for title, value in stats: - draw.rounded_rectangle((sx, 1170, sx + 320, 1262), radius=18, fill=PANEL, outline=(35, 49, 59), width=1) + draw.rounded_rectangle( + (sx, 1170, sx + 320, 1262), + radius=18, + fill=PANEL, + outline=(35, 49, 59), + width=1, + ) draw.text((sx + 22, 1184), title, fill=MUTED, font=ui_font(16)) draw.text((sx + 22, 1208), value, fill=INK, font=ui_font(34, True)) sx += 344 - draw.text((sx + 20, 1196), "solid = text carrier dashed = image carrier thin rungs = pair gap", fill=MUTED, font=ui_font(18)) + draw.text( + (sx + 20, 1196), + "solid = text carrier dashed = image carrier thin rungs = pair gap", + fill=MUTED, + font=ui_font(18), + ) out = Path(args.out) out.parent.mkdir(parents=True, exist_ok=True) diff --git a/packages/snapcompact/research/snapcompact_convergence_extras.py b/packages/snapcompact/research/snapcompact_convergence_extras.py index 7ab101d6b..f516e1bef 100644 --- a/packages/snapcompact/research/snapcompact_convergence_extras.py +++ b/packages/snapcompact/research/snapcompact_convergence_extras.py @@ -31,8 +31,12 @@ PALETTE = { def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: for path in [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ]: if Path(path).exists(): return ImageFont.truetype(path, size) @@ -69,43 +73,87 @@ def background(w: int, h: int) -> Image.Image: gd = ImageDraw.Draw(glow) gd.ellipse((-240, -200, 880, 680), fill=(75, 220, 255, 25)) gd.ellipse((w - 1000, h - 760, w + 240, h + 220), fill=(255, 112, 72, 25)) - return Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(84))).convert("RGB") + return Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(84)) + ).convert("RGB") -def render_funnel(out_path: Path, text_arr: np.ndarray, image_arr: np.ndarray, layers_meta: list[dict[str, Any]], best_layer: int, records: list[dict[str, Any]]) -> None: +def render_funnel( + out_path: Path, + text_arr: np.ndarray, + image_arr: np.ndarray, + layers_meta: list[dict[str, Any]], + best_layer: int, + records: list[dict[str, Any]], +) -> None: n_q, n_layers, _ = text_arr.shape snapshots = [1, max(2, best_layer // 2), best_layer] # Shared PCA frame from the peak layer keeps the panels comparable. - ref = np.concatenate([center(text_arr[:, best_layer, :]), center(image_arr[:, best_layer, :])], axis=0) + ref = np.concatenate( + [center(text_arr[:, best_layer, :]), center(image_arr[:, best_layer, :])], + axis=0, + ) _, _, vt = np.linalg.svd(ref, full_matrices=False) basis = vt[:2].T # [D, 2] w, h = 2200, 1240 canvas = background(w, h) draw = ImageDraw.Draw(canvas) - draw.text((64, 42), "QWEN CARRIER CONVERGENCE — TRAJECTORY VIEW", fill=PALETTE["amber"], font=ui_font(24, True)) - draw.text((64, 84), "Watch the two carriers fuse", fill=PALETTE["ink"], font=ui_font(64, True)) - draw.text((66, 164), "Each color is one question; ● came in as text, ◆ came in as pixels. Same 2D projection at every depth. The tie-lines shrink as carriers converge.", fill=PALETTE["muted"], font=ui_font(23)) + draw.text( + (64, 42), + "QWEN CARRIER CONVERGENCE — TRAJECTORY VIEW", + fill=PALETTE["amber"], + font=ui_font(24, True), + ) + draw.text( + (64, 84), + "Watch the two carriers fuse", + fill=PALETTE["ink"], + font=ui_font(64, True), + ) + draw.text( + (66, 164), + "Each color is one question; ● came in as text, ◆ came in as pixels. Same 2D projection at every depth. The tie-lines shrink as carriers converge.", + fill=PALETTE["muted"], + font=ui_font(23), + ) panel_w = 660 titles = ["early (layer {})", "middle (layer {})", "peak (layer {})"] for pi, (layer, title) in enumerate(zip(snapshots, titles)): x0 = 64 + pi * (panel_w + 44) box = (x0, 232, x0 + panel_w, 952) - draw.rounded_rectangle(box, radius=24, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((x0 + 26, 252), title.format(layer), fill=PALETTE["ink"], font=ui_font(27, True)) + draw.rounded_rectangle( + box, radius=24, fill=PALETTE["panel"], outline=(35, 49, 59), width=1 + ) + draw.text( + (x0 + 26, 252), + title.format(layer), + fill=PALETTE["ink"], + font=ui_font(27, True), + ) t_proj = center(text_arr[:, layer, :]) @ basis i_proj = center(image_arr[:, layer, :]) @ basis both = np.concatenate([t_proj, i_proj], axis=0) lim = float(np.abs(both).max()) * 1.15 or 1.0 gx0, gy0, gx1, gy1 = x0 + 36, 306, x0 + panel_w - 36, 912 + def to_px(p: np.ndarray) -> tuple[int, int]: return ( round(gx0 + (p[0] + lim) / (2 * lim) * (gx1 - gx0)), round(gy0 + (1 - (p[1] + lim) / (2 * lim)) * (gy1 - gy0)), ) - draw.line((gx0, (gy0 + gy1) // 2, gx1, (gy0 + gy1) // 2), fill=PALETTE["grid"], width=1) - draw.line(((gx0 + gx1) // 2, gy0, (gx0 + gx1) // 2, gy1), fill=PALETTE["grid"], width=1) + + draw.line( + (gx0, (gy0 + gy1) // 2, gx1, (gy0 + gy1) // 2), + fill=PALETTE["grid"], + width=1, + ) + draw.line( + ((gx0 + gx1) // 2, gy0, (gx0 + gx1) // 2, gy1), + fill=PALETTE["grid"], + width=1, + ) pair_dist = 0.0 for qi in range(n_q): color = question_hue(qi, n_q) @@ -113,19 +161,45 @@ def render_funnel(out_path: Path, text_arr: np.ndarray, image_arr: np.ndarray, l ip = to_px(i_proj[qi]) draw.line((tp, ip), fill=(*color, 0)[:3], width=3) r = 11 - draw.ellipse((tp[0] - r, tp[1] - r, tp[0] + r, tp[1] + r), fill=color, outline=(8, 10, 12), width=2) + draw.ellipse( + (tp[0] - r, tp[1] - r, tp[0] + r, tp[1] + r), + fill=color, + outline=(8, 10, 12), + width=2, + ) d = ImageDraw.Draw(canvas) - d.polygon([(ip[0], ip[1] - r - 2), (ip[0] + r + 2, ip[1]), (ip[0], ip[1] + r + 2), (ip[0] - r - 2, ip[1])], fill=color, outline=(8, 10, 12)) + d.polygon( + [ + (ip[0], ip[1] - r - 2), + (ip[0] + r + 2, ip[1]), + (ip[0], ip[1] + r + 2), + (ip[0] - r - 2, ip[1]), + ], + fill=color, + outline=(8, 10, 12), + ) pair_dist += float(np.linalg.norm(t_proj[qi] - i_proj[qi])) pair_dist /= n_q norm_dist = pair_dist / (2 * lim) meta = layers_meta[layer] - draw.text((x0 + 26, 916), f"mean pair gap: {norm_dist * 100:.0f}% of frame · matched cos {meta['matched_cosine']:.2f}", fill=PALETTE["muted"], font=ui_font(17)) + draw.text( + (x0 + 26, 916), + f"mean pair gap: {norm_dist * 100:.0f}% of frame · matched cos {meta['matched_cosine']:.2f}", + fill=PALETTE["muted"], + font=ui_font(17), + ) # Pair-distance by layer strip. strip = (64, 996, 2136, 1190) - draw.rounded_rectangle(strip, radius=24, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((96, 1014), "matched-pair separation by layer (lower = carriers agree)", fill=PALETTE["ink"], font=ui_font(22, True)) + draw.rounded_rectangle( + strip, radius=24, fill=PALETTE["panel"], outline=(35, 49, 59), width=1 + ) + draw.text( + (96, 1014), + "matched-pair separation by layer (lower = carriers agree)", + fill=PALETTE["ink"], + font=ui_font(22, True), + ) gx0, gy0, gx1, gy1 = 110, 1062, 2100, 1162 gaps = [] for layer in range(n_layers): @@ -141,15 +215,24 @@ def render_funnel(out_path: Path, text_arr: np.ndarray, image_arr: np.ndarray, l xb = gx0 + (layer + 1) * bw - 3 bh = (gy1 - gy0) * gap / hi color = PALETTE["orange"] if layer == best_layer else (62, 86, 102) - draw.rounded_rectangle((round(xa), round(gy1 - bh), round(xb), gy1), radius=5, fill=color) + draw.rounded_rectangle( + (round(xa), round(gy1 - bh), round(xb), gy1), radius=5, fill=color + ) draw.text((gx0, gy1 + 6), "layer 0", fill=PALETTE["muted"], font=ui_font(13)) - draw.text((gx1 - 70, gy1 + 6), f"layer {n_layers - 1}", fill=PALETTE["muted"], font=ui_font(13)) + draw.text( + (gx1 - 70, gy1 + 6), + f"layer {n_layers - 1}", + fill=PALETTE["muted"], + font=ui_font(13), + ) out_path.parent.mkdir(parents=True, exist_ok=True) canvas.save(out_path) -def render_gif(out_path: Path, cross_sim: np.ndarray, layers_meta: list[dict[str, Any]]) -> None: +def render_gif( + out_path: Path, cross_sim: np.ndarray, layers_meta: list[dict[str, Any]] +) -> None: n_layers, n_q, _ = cross_sim.shape cell = 46 pad = 36 @@ -162,29 +245,58 @@ def render_gif(out_path: Path, cross_sim: np.ndarray, layers_meta: list[dict[str draw = ImageDraw.Draw(frame) for y in range(0, h, 14): draw.line((0, y, w, y), fill=(7, 10 + y % 9, 15 + y % 11)) - draw.text((pad, 22), "cross-carrier matching", fill=PALETTE["ink"], font=ui_font(30, True)) - draw.text((pad, 62), "text question i × image question j", fill=PALETTE["muted"], font=ui_font(17)) + draw.text( + (pad, 22), + "cross-carrier matching", + fill=PALETTE["ink"], + font=ui_font(30, True), + ) + draw.text( + (pad, 62), + "text question i × image question j", + fill=PALETTE["muted"], + font=ui_font(17), + ) meta = layers_meta[layer] - draw.text((pad, 92), f"layer {layer:02d} matched {meta['matched_cosine']:+.2f} others {meta['mismatched_cosine']:+.2f}", fill=PALETTE["amber"], font=ui_font(19, True)) + draw.text( + (pad, 92), + f"layer {layer:02d} matched {meta['matched_cosine']:+.2f} others {meta['mismatched_cosine']:+.2f}", + fill=PALETTE["amber"], + font=ui_font(19, True), + ) for r in range(n_q): for c in range(n_q): xa = pad + c * cell ya = header + r * cell - draw.rounded_rectangle((xa, ya, xa + cell - 4, ya + cell - 4), radius=7, fill=diverging_color(float(cross_sim[layer, r, c]))) + draw.rounded_rectangle( + (xa, ya, xa + cell - 4, ya + cell - 4), + radius=7, + fill=diverging_color(float(cross_sim[layer, r, c])), + ) # progress bar bar_y = header + n_q * cell + 18 - draw.rounded_rectangle((pad, bar_y, w - pad, bar_y + 10), radius=5, fill=(30, 40, 48)) - draw.rounded_rectangle((pad, bar_y, pad + (w - 2 * pad) * (layer + 1) // n_layers, bar_y + 10), radius=5, fill=PALETTE["cyan"]) + draw.rounded_rectangle( + (pad, bar_y, w - pad, bar_y + 10), radius=5, fill=(30, 40, 48) + ) + draw.rounded_rectangle( + (pad, bar_y, pad + (w - 2 * pad) * (layer + 1) // n_layers, bar_y + 10), + radius=5, + fill=PALETTE["cyan"], + ) frames.append(frame) durations = [240] * n_layers durations[-1] = 2200 out_path.parent.mkdir(parents=True, exist_ok=True) - frames[0].save(out_path, save_all=True, append_images=frames[1:], duration=durations, loop=0) + frames[0].save( + out_path, save_all=True, append_images=frames[1:], duration=durations, loop=0 + ) def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--result-dir", default=str(HERE / "results" / "qwen-carrier-convergence-n12")) + ap.add_argument( + "--result-dir", default=str(HERE / "results" / "qwen-carrier-convergence-n12") + ) args = ap.parse_args() result_dir = Path(args.result_dir) summary = json.loads((result_dir / "summary.json").read_text()) @@ -197,7 +309,9 @@ def main() -> None: funnel_path = result_dir / "convergence-funnel.png" gif_path = result_dir / "diagonal-emerges.gif" - render_funnel(funnel_path, text_arr, image_arr, layers_meta, best_layer, summary["records"]) + render_funnel( + funnel_path, text_arr, image_arr, layers_meta, best_layer, summary["records"] + ) render_gif(gif_path, cross_sim, layers_meta) print(funnel_path) print(gif_path) diff --git a/packages/snapcompact/research/snapcompact_convergence_viz.py b/packages/snapcompact/research/snapcompact_convergence_viz.py index b28b33088..6a1ea5624 100644 --- a/packages/snapcompact/research/snapcompact_convergence_viz.py +++ b/packages/snapcompact/research/snapcompact_convergence_viz.py @@ -32,8 +32,12 @@ PALETTE = { def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: for path in [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ]: if Path(path).exists(): return ImageFont.truetype(path, size) @@ -41,7 +45,10 @@ def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: def mono_font(size: int) -> ImageFont.ImageFont: - for path in ["/System/Library/Fonts/Monaco.ttf", "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf"]: + for path in [ + "/System/Library/Fonts/Monaco.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf", + ]: if Path(path).exists(): return ImageFont.truetype(path, size) return ImageFont.load_default() @@ -57,9 +64,19 @@ def diverging_color(t: float) -> tuple[int, int, int]: return (round(8 + 247 * u), round(20 + 130 * u), round(34 + 20 * u)) -def draw_matrix(draw: ImageDraw.ImageDraw, mat: np.ndarray, box: tuple[int, int, int, int], title: str, subtitle: str, color: tuple[int, int, int], highlight_diag: bool = False) -> None: +def draw_matrix( + draw: ImageDraw.ImageDraw, + mat: np.ndarray, + box: tuple[int, int, int, int], + title: str, + subtitle: str, + color: tuple[int, int, int], + highlight_diag: bool = False, +) -> None: x0, y0, x1, y1 = box - draw.rounded_rectangle(box, radius=20, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1) + draw.rounded_rectangle( + box, radius=20, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1 + ) draw.text((x0 + 20, y0 + 16), title, fill=color, font=ui_font(24, True)) draw.text((x0 + 20, y0 + 48), subtitle, fill=PALETTE["muted"], font=ui_font(15)) n = mat.shape[0] @@ -72,27 +89,57 @@ def draw_matrix(draw: ImageDraw.ImageDraw, mat: np.ndarray, box: tuple[int, int, xb = round(gx0 + (c + 1) * cell) - 2 ya = round(gy0 + r * cell) yb = round(gy0 + (r + 1) * cell) - 2 - draw.rounded_rectangle((xa, ya, xb, yb), radius=4, fill=diverging_color(float(mat[r, c]))) + draw.rounded_rectangle( + (xa, ya, xb, yb), radius=4, fill=diverging_color(float(mat[r, c])) + ) if highlight_diag: for r in range(n): xa = round(gx0 + r * cell) ya = round(gy0 + r * cell) - draw.rounded_rectangle((xa - 1, ya - 1, round(xa + cell) - 1, round(ya + cell) - 1), radius=5, outline=PALETTE["amber"], width=2) - draw.text((gx0, round(gy0 + side) + 6), "questions →", fill=PALETTE["muted"], font=ui_font(13)) + draw.rounded_rectangle( + (xa - 1, ya - 1, round(xa + cell) - 1, round(ya + cell) - 1), + radius=5, + outline=PALETTE["amber"], + width=2, + ) + draw.text( + (gx0, round(gy0 + side) + 6), + "questions →", + fill=PALETTE["muted"], + font=ui_font(13), + ) -def draw_curves(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], layers: list[dict[str, Any]]) -> None: +def draw_curves( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + layers: list[dict[str, Any]], +) -> None: x0, y0, x1, y1 = box - draw.rounded_rectangle(box, radius=20, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1) - draw.text((x0 + 22, y0 + 16), "convergence by depth", fill=PALETTE["ink"], font=ui_font(24, True)) - draw.text((x0 + 22, y0 + 48), "carrier-centered cosine: same question across carriers vs different questions", fill=PALETTE["muted"], font=ui_font(15)) + draw.rounded_rectangle( + box, radius=20, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1 + ) + draw.text( + (x0 + 22, y0 + 16), + "convergence by depth", + fill=PALETTE["ink"], + font=ui_font(24, True), + ) + draw.text( + (x0 + 22, y0 + 48), + "carrier-centered cosine: same question across carriers vs different questions", + fill=PALETTE["muted"], + font=ui_font(15), + ) gx0, gy0, gx1, gy1 = x0 + 52, y0 + 92, x1 - 26, y1 - 56 lo, hi = -0.15, 1.0 for i in range(5): y = gy0 + (gy1 - gy0) * i / 4 draw.line((gx0, y, gx1, y), fill=PALETTE["grid"], width=1) value = hi - (hi - lo) * i / 4 - draw.text((x0 + 12, y - 8), f"{value:.1f}", fill=PALETTE["muted"], font=ui_font(12)) + draw.text( + (x0 + 12, y - 8), f"{value:.1f}", fill=PALETTE["muted"], font=ui_font(12) + ) series = [ ("matched_cosine", PALETTE["amber"], 6), ("mismatched_cosine", PALETTE["muted"], 4), @@ -113,8 +160,14 @@ def draw_curves(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], layer continue draw.line(pts, fill=color, width=width, joint="curve") draw.text((gx0, gy1 + 14), "layer 0", fill=PALETTE["muted"], font=ui_font(13)) - draw.text((gx1 - 64, gy1 + 14), f"layer {n - 1}", fill=PALETTE["muted"], font=ui_font(13)) - legend = [("same question, text↔image", PALETTE["amber"]), ("different questions", PALETTE["muted"]), ("RSA geometry corr", PALETTE["cyan"])] + draw.text( + (gx1 - 64, gy1 + 14), f"layer {n - 1}", fill=PALETTE["muted"], font=ui_font(13) + ) + legend = [ + ("same question, text↔image", PALETTE["amber"]), + ("different questions", PALETTE["muted"]), + ("RSA geometry corr", PALETTE["cyan"]), + ] lx = gx0 for label, color in legend: draw.rounded_rectangle((lx, y0 + 70, lx + 18, y0 + 78), radius=4, fill=color) @@ -122,12 +175,33 @@ def draw_curves(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], layer lx += 232 -def draw_answers(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], records: list[dict[str, Any]]) -> None: +def draw_answers( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + records: list[dict[str, Any]], +) -> None: x0, y0, x1, y1 = box - draw.rounded_rectangle(box, radius=20, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1) - draw.text((x0 + 22, y0 + 16), "behavioral check: both carriers answer alike", fill=PALETTE["ink"], font=ui_font(24, True)) - draw.text((x0 + 240, y0 + 56), "text carrier", fill=PALETTE["cyan"], font=ui_font(15, True)) - draw.text((x0 + 470, y0 + 56), "image carrier", fill=PALETTE["orange"], font=ui_font(15, True)) + draw.rounded_rectangle( + box, radius=20, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1 + ) + draw.text( + (x0 + 22, y0 + 16), + "behavioral check: both carriers answer alike", + fill=PALETTE["ink"], + font=ui_font(24, True), + ) + draw.text( + (x0 + 240, y0 + 56), + "text carrier", + fill=PALETTE["cyan"], + font=ui_font(15, True), + ) + draw.text( + (x0 + 470, y0 + 56), + "image carrier", + fill=PALETTE["orange"], + font=ui_font(15, True), + ) y = y0 + 84 row_h = (y1 - y0 - 96) // len(records) fnt = mono_font(15) @@ -137,14 +211,29 @@ def draw_answers(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], reco draw.text((x0 + 240, y), r["text_answer"][:22], fill=PALETTE["ink"], font=fnt) draw.text((x0 + 470, y), r["image_answer"][:22], fill=PALETTE["ink"], font=fnt) mark = "=" if r["agree"] else "≠" - draw.text((x1 - 44, y), mark, fill=PALETTE["green"] if r["agree"] else PALETTE["red"], font=ui_font(17, True)) + draw.text( + (x1 - 44, y), + mark, + fill=PALETTE["green"] if r["agree"] else PALETTE["red"], + font=ui_font(17, True), + ) y += row_h def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--result-dir", default=str(HERE / "results" / "qwen-carrier-convergence-n12")) - ap.add_argument("--out", default=str(HERE / "results" / "qwen-carrier-convergence-n12" / "carrier-convergence.png")) + ap.add_argument( + "--result-dir", default=str(HERE / "results" / "qwen-carrier-convergence-n12") + ) + ap.add_argument( + "--out", + default=str( + HERE + / "results" + / "qwen-carrier-convergence-n12" + / "carrier-convergence.png" + ), + ) args = ap.parse_args() result_dir = Path(args.result_dir) @@ -166,12 +255,24 @@ def main() -> None: gd = ImageDraw.Draw(glow) gd.ellipse((-240, -200, 900, 700), fill=(75, 220, 255, 26)) gd.ellipse((1300, 180, 2480, 1380), fill=(255, 112, 72, 24)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(86))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(86)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) best = summary["best"] - draw.text((64, 42), "QWEN CARRIER CONVERGENCE", fill=PALETTE["amber"], font=ui_font(24, True)) - draw.text((64, 84), "Two carriers, one thought", fill=PALETTE["ink"], font=ui_font(66, True)) + draw.text( + (64, 42), + "QWEN CARRIER CONVERGENCE", + fill=PALETTE["amber"], + font=ui_font(24, True), + ) + draw.text( + (64, 84), + "Two carriers, one thought", + fill=PALETTE["ink"], + font=ui_font(66, True), + ) draw.text( (66, 166), "Hidden state at the answer position, carrier means removed. Same question through text or bitmap lands in the same place; different questions do not.", @@ -183,21 +284,57 @@ def main() -> None: ("matched pairs", f"{best['matched_cosine']:.2f}", "same Q, text ↔ image"), ("mismatched pairs", f"{best['mismatched_cosine']:.2f}", "different questions"), ("RSA geometry corr", f"{best['rsa_pearson']:.2f}", f"layer {best['layer']}"), - ("pair retrieval", f"{best['match_rank_accuracy'] * 100:.0f}%", "nearest cross-carrier match"), - ("answer agreement", f"{summary['answer_agreement'] * 100:.0f}%", "text vs image generations"), + ( + "pair retrieval", + f"{best['match_rank_accuracy'] * 100:.0f}%", + "nearest cross-carrier match", + ), + ( + "answer agreement", + f"{summary['answer_agreement'] * 100:.0f}%", + "text vs image generations", + ), ] sx = 64 for title, value, caption in stats: - draw.rounded_rectangle((sx, 222, sx + 396, 332), radius=20, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) + draw.rounded_rectangle( + (sx, 222, sx + 396, 332), + radius=20, + fill=PALETTE["panel"], + outline=(35, 49, 59), + width=1, + ) draw.text((sx + 22, 240), title, fill=PALETTE["muted"], font=ui_font(16)) draw.text((sx + 22, 264), value, fill=PALETTE["ink"], font=ui_font(40, True)) draw.text((sx + 226, 290), caption, fill=PALETTE["muted"], font=ui_font(13)) sx += 420 n = text_sim.shape[0] - draw_matrix(draw, text_sim, (64, 376, 600, 952), "text-carrier geometry", f"{n}×{n} question similarity, layer {best_layer}", PALETTE["cyan"]) - draw_matrix(draw, image_sim, (628, 376, 1164, 952), "image-carrier geometry", "same questions through the bitmap — same shape", PALETTE["orange"]) - draw_matrix(draw, cross_sim, (1192, 376, 1728, 952), "cross-carrier matching", "text question i × image question j — bright diagonal", PALETTE["green"], highlight_diag=True) + draw_matrix( + draw, + text_sim, + (64, 376, 600, 952), + "text-carrier geometry", + f"{n}×{n} question similarity, layer {best_layer}", + PALETTE["cyan"], + ) + draw_matrix( + draw, + image_sim, + (628, 376, 1164, 952), + "image-carrier geometry", + "same questions through the bitmap — same shape", + PALETTE["orange"], + ) + draw_matrix( + draw, + cross_sim, + (1192, 376, 1728, 952), + "cross-carrier matching", + "text question i × image question j — bright diagonal", + PALETTE["green"], + highlight_diag=True, + ) draw_curves(draw, (64, 996, 1164, 1264), layers) draw_answers(draw, (1192, 996, 2136, 1264), records) diff --git a/packages/snapcompact/research/snapcompact_lockon_anatomy_viz.py b/packages/snapcompact/research/snapcompact_lockon_anatomy_viz.py index 357c63173..c22fd817f 100644 --- a/packages/snapcompact/research/snapcompact_lockon_anatomy_viz.py +++ b/packages/snapcompact/research/snapcompact_lockon_anatomy_viz.py @@ -47,8 +47,12 @@ COND_COLORS = { def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: for path in [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ]: if Path(path).exists(): return ImageFont.truetype(path, size) @@ -56,7 +60,10 @@ def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: def mono_font(size: int) -> ImageFont.ImageFont: - for path in ["/System/Library/Fonts/Monaco.ttf", "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf"]: + for path in [ + "/System/Library/Fonts/Monaco.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf", + ]: if Path(path).exists(): return ImageFont.truetype(path, size) return ImageFont.load_default() @@ -72,18 +79,43 @@ def label_font(label: str, size: int) -> ImageFont.ImageFont: return mono_font(size) -def crosshair(draw: ImageDraw.ImageDraw, cx: int, cy: int, r: int, color: tuple[int, int, int], width: int = 4) -> None: +def crosshair( + draw: ImageDraw.ImageDraw, + cx: int, + cy: int, + r: int, + color: tuple[int, int, int], + width: int = 4, +) -> None: draw.ellipse((cx - r, cy - r, cx + r, cy + r), outline=color, width=width) - draw.ellipse((cx - r // 2, cy - r // 2, cx + r // 2, cy + r // 2), outline=color, width=2) + draw.ellipse( + (cx - r // 2, cy - r // 2, cx + r // 2, cy + r // 2), outline=color, width=2 + ) for dx, dy in ((-1, 0), (1, 0), (0, -1), (0, 1)): - draw.line((cx + dx * (r - 6), cy + dy * (r - 6), cx + dx * (r + 14), cy + dy * (r + 14)), fill=color, width=width) + draw.line( + ( + cx + dx * (r - 6), + cy + dy * (r - 6), + cx + dx * (r + 14), + cy + dy * (r + 14), + ), + fill=color, + width=width, + ) def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--result-dir", default=str(HERE / "results" / "qwen-materialize-sweep-q3")) + ap.add_argument( + "--result-dir", default=str(HERE / "results" / "qwen-materialize-sweep-q3") + ) ap.add_argument("--condition", default="base-8x13") - ap.add_argument("--out", default=str(HERE / "results" / "qwen-materialize-sweep-q3" / "lockon-anatomy.png")) + ap.add_argument( + "--out", + default=str( + HERE / "results" / "qwen-materialize-sweep-q3" / "lockon-anatomy.png" + ), + ) args = ap.parse_args() result_dir = Path(args.result_dir) summary = json.loads((result_dir / "summary.json").read_text()) @@ -106,11 +138,20 @@ def main() -> None: gd.ellipse((520, 620, 1280, 1180), fill=(255, 196, 68, 36)) gd.ellipse((-260, -240, 760, 560), fill=(75, 220, 255, 26)) gd.ellipse((1500, -100, 2480, 700), fill=(255, 112, 72, 20)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(90))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(90)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) - draw.text((64, 40), "THE LOCK-ON INSTRUMENT", fill=P["amber"], font=ui_font(24, True)) - draw.text((64, 80), "How we decide where the answer materializes", fill=P["ink"], font=ui_font(58, True)) + draw.text( + (64, 40), "THE LOCK-ON INSTRUMENT", fill=P["amber"], font=ui_font(24, True) + ) + draw.text( + (64, 80), + "How we decide where the answer materializes", + fill=P["ink"], + font=ui_font(58, True), + ) draw.text( (66, 154), "At every layer, a logit-lens probe taps the answer patch's residual stream: final RMSNorm → LM head → softmax over 152k vocabulary entries.", @@ -126,23 +167,39 @@ def main() -> None: # ---- Probe pipeline card (top left). pipe = (64, 248, 700, 420) - draw.rounded_rectangle(pipe, radius=22, fill=P["panel"], outline=(35, 49, 59), width=1) - draw.text((92, 268), "the probe, applied at every layer ℓ", fill=P["ink"], font=ui_font(22, True)) + draw.rounded_rectangle( + pipe, radius=22, fill=P["panel"], outline=(35, 49, 59), width=1 + ) + draw.text( + (92, 268), + "the probe, applied at every layer ℓ", + fill=P["ink"], + font=ui_font(22, True), + ) stages = ["h(patch)", "RMSNorm", "LM head", "softmax", "top-1?"] sx = 92 for si, stage in enumerate(stages): color = P["amber"] if si == len(stages) - 1 else P["cyan"] tw = int(draw.textlength(stage, font=mono_font(16))) + 24 - draw.rounded_rectangle((sx, 318, sx + tw, 356), radius=10, fill=P["panel2"], outline=color, width=2) + draw.rounded_rectangle( + (sx, 318, sx + tw, 356), radius=10, fill=P["panel2"], outline=color, width=2 + ) draw.text((sx + 12, 327), stage, fill=color, font=mono_font(16)) if si < len(stages) - 1: draw.text((sx + tw + 4, 327), "→", fill=P["faint"], font=ui_font(18)) sx += tw + 28 - draw.text((92, 376), f"vocabulary = 152k entries · answer BPEs = {answer_strs}", fill=P["muted"], font=mono_font(14)) + draw.text( + (92, 376), + f"vocabulary = 152k entries · answer BPEs = {answer_strs}", + fill=P["muted"], + font=mono_font(14), + ) # ---- The patch under test (left). patch_card = (64, 460, 380, 760) - draw.rounded_rectangle(patch_card, radius=22, fill=P["panel"], outline=(35, 49, 59), width=1) + draw.rounded_rectangle( + patch_card, radius=22, fill=P["panel"], outline=(35, 49, 59), width=1 + ) draw.text((92, 480), "specimen", fill=P["orange"], font=ui_font(21, True)) carrier = Image.open(result_dir / "images" / f"{args.condition}.png").convert("RGB") rw = 1568 @@ -151,35 +208,78 @@ def main() -> None: lock_entry = layers[lock_on] tok_idx = lock_entry["best_token_index"] r0, c0 = tok_idx // grid, tok_idx % grid - cell = carrier.resize((rw, rw), Image.Resampling.LANCZOS).crop((c0 * px, r0 * px, (c0 + 1) * px, (r0 + 1) * px)) + cell = carrier.resize((rw, rw), Image.Resampling.LANCZOS).crop( + (c0 * px, r0 * px, (c0 + 1) * px, (r0 + 1) * px) + ) big = cell.resize((196, 196), Image.Resampling.NEAREST) - draw.rounded_rectangle((118, 516, 326, 724), radius=12, fill=(244, 242, 230), outline=P["orange"], width=4) + draw.rounded_rectangle( + (118, 516, 326, 724), + radius=12, + fill=(244, 242, 230), + outline=P["orange"], + width=4, + ) canvas.paste(big, (124, 522)) - draw.text((118, 730), f"visual token #{tok_idx} · 28×28 px", fill=P["muted"], font=mono_font(13)) + draw.text( + (118, 730), + f"visual token #{tok_idx} · 28×28 px", + fill=P["muted"], + font=mono_font(13), + ) # ---- Depth shaft. shaft_x = 470 shaft_top, shaft_bot = 470, 1340 - draw.rounded_rectangle((shaft_x - 7, shaft_top, shaft_x + 7, shaft_bot), radius=7, fill=(20, 28, 35), outline=(40, 54, 64), width=1) + draw.rounded_rectangle( + (shaft_x - 7, shaft_top, shaft_x + 7, shaft_bot), + radius=7, + fill=(20, 28, 35), + outline=(40, 54, 64), + width=1, + ) def layer_y(layer: int) -> int: return round(shaft_top + (shaft_bot - shaft_top) * layer / (n_layers - 1)) # p(answer) trajectory along the shaft. - traj = [(shaft_x + 14 + 230 * min(1.0, e["best_answer_p"]), layer_y(e["layer"])) for e in layers] + traj = [ + (shaft_x + 14 + 230 * min(1.0, e["best_answer_p"]), layer_y(e["layer"])) + for e in layers + ] for i in range(len(traj) - 1): draw.line((traj[i], traj[i + 1]), fill=(120, 96, 40), width=3) - draw.text((shaft_x + 30, shaft_bot + 10), "p(answer BPE) →", fill=(150, 124, 60), font=ui_font(14)) + draw.text( + (shaft_x + 30, shaft_bot + 10), + "p(answer BPE) →", + fill=(150, 124, 60), + font=ui_font(14), + ) for layer in range(n_layers): y = layer_y(layer) major = layer % 4 == 0 or layer == n_layers - 1 - draw.line((shaft_x - (16 if major else 10), y, shaft_x + (16 if major else 10), y), fill=P["faint"] if major else (52, 64, 73), width=2) + draw.line( + (shaft_x - (16 if major else 10), y, shaft_x + (16 if major else 10), y), + fill=P["faint"] if major else (52, 64, 73), + width=2, + ) if major: - draw.text((shaft_x - 58, y - 9), f"L{layer:02d}", fill=P["muted"], font=mono_font(13)) + draw.text( + (shaft_x - 58, y - 9), + f"L{layer:02d}", + fill=P["muted"], + font=mono_font(13), + ) # Patch entering the shaft. draw.line((326, 620, shaft_x - 18, shaft_top + 6), fill=P["orange"], width=3) - draw.polygon([(shaft_x - 14, shaft_top + 2), (shaft_x - 30, shaft_top - 4), (shaft_x - 26, shaft_top + 16)], fill=P["orange"]) + draw.polygon( + [ + (shaft_x - 14, shaft_top + 2), + (shaft_x - 30, shaft_top - 4), + (shaft_x - 26, shaft_top + 16), + ], + fill=P["orange"], + ) # ---- Readout cards at sampled depths (real top-5). samples = [0, 10, 18, lock_on, n_layers - 1] @@ -199,16 +299,36 @@ def main() -> None: for layer, cy in zip(samples, card_ys): entry = layers[layer] is_lock = layer == lock_on - accent = P["amber"] if is_lock else P["cyan"] if entry["best_answer_p"] > 0.01 else P["faint"] + accent = ( + P["amber"] + if is_lock + else P["cyan"] + if entry["best_answer_p"] > 0.01 + else P["faint"] + ) # Connector. ly = layer_y(layer) - draw.line((shaft_x + 16, ly, card_x - 18, cy + card_h // 2), fill=accent, width=3 if is_lock else 2) + draw.line( + (shaft_x + 16, ly, card_x - 18, cy + card_h // 2), + fill=accent, + width=3 if is_lock else 2, + ) draw.ellipse((shaft_x + 12, ly - 5, shaft_x + 22, ly + 5), fill=accent) - draw.rounded_rectangle((card_x, cy, card_x + card_w, cy + card_h), radius=16, fill=P["panel2"], outline=accent, width=3 if is_lock else 1) + draw.rounded_rectangle( + (card_x, cy, card_x + card_w, cy + card_h), + radius=16, + fill=P["panel2"], + outline=accent, + width=3 if is_lock else 1, + ) title = f"L{layer:02d} readout" + (" LOCK-ON" if is_lock else "") draw.text((card_x + 20, cy + 10), title, fill=accent, font=ui_font(19, True)) if is_lock: - tx = card_x + 20 + draw.textlength(f"L{layer:02d} readout ", font=ui_font(19, True)) + tx = ( + card_x + + 20 + + draw.textlength(f"L{layer:02d} readout ", font=ui_font(19, True)) + ) draw.ellipse((tx - 8, cy + 14, tx + 4, cy + 26), outline=accent, width=3) bx = card_x + 20 by = cy + 44 @@ -220,20 +340,52 @@ def main() -> None: pill_w = 108 fill = (66, 92, 36) if hit else (16, 22, 28) outline = P["green"] if hit else (38, 52, 61) - draw.rounded_rectangle((bx, by, bx + pill_w, by + 30), radius=8, fill=fill, outline=outline, width=2) - draw.text((bx + 8, by + 6), label, fill=(220, 255, 190) if hit else P["ink"], font=label_font(label, 13)) + draw.rounded_rectangle( + (bx, by, bx + pill_w, by + 30), + radius=8, + fill=fill, + outline=outline, + width=2, + ) + draw.text( + (bx + 8, by + 6), + label, + fill=(220, 255, 190) if hit else P["ink"], + font=label_font(label, 13), + ) bar = round(min(1.0, t["p"] / 0.4) * pill_w) - draw.rounded_rectangle((bx, by + 36, bx + max(3, bar), by + 42), radius=3, fill=P["amber"] if hit else (60, 76, 88)) - draw.text((bx, by + 46, ), f"{t['p']:.3f}", fill=P["muted"], font=mono_font(10)) + draw.rounded_rectangle( + (bx, by + 36, bx + max(3, bar), by + 42), + radius=3, + fill=P["amber"] if hit else (60, 76, 88), + ) + draw.text( + ( + bx, + by + 46, + ), + f"{t['p']:.3f}", + fill=P["muted"], + font=mono_font(10), + ) bx += pill_w + 12 if is_lock: crosshair(draw, shaft_x, ly, 26, P["amber"], 4) - draw.text((shaft_x + 44, ly + 26), f"first top-1 hit: “{entry['best_token_top'][0]['str'].strip()}” p={entry['best_token_top'][0]['p']:.2f}", fill=P["amber"], font=ui_font(16, True)) + draw.text( + (shaft_x + 44, ly + 26), + f"first top-1 hit: “{entry['best_token_top'][0]['str'].strip()}” p={entry['best_token_top'][0]['p']:.2f}", + fill=P["amber"], + font=ui_font(16, True), + ) # ---- Why it matters (right column). why = (1460, 248, 2136, 716) - draw.rounded_rectangle(why, radius=22, fill=P["panel"], outline=(35, 49, 59), width=1) - draw.text((1492, 270), "why lock-on is the metric", fill=P["ink"], font=ui_font(26, True)) + draw.rounded_rectangle( + why, radius=22, fill=P["panel"], outline=(35, 49, 59), width=1 + ) + draw.text( + (1492, 270), "why lock-on is the metric", fill=P["ink"], font=ui_font(26, True) + ) lines = [ ("It separates decoding from reasoning.", P["ink"]), ("Layers before lock-on are spent turning", P["muted"]), @@ -247,7 +399,12 @@ def main() -> None: ty += 30 draw.line((1492, ty + 8, 2104, ty + 8), fill=P["grid"], width=1) ty += 26 - draw.text((1492, ty), "reasoning budget after lock-on", fill=P["muted"], font=ui_font(16, True)) + draw.text( + (1492, ty), + "reasoning budget after lock-on", + fill=P["muted"], + font=ui_font(16, True), + ) ty += 30 for name, color in COND_COLORS.items(): c = conditions.get(name) @@ -256,13 +413,22 @@ def main() -> None: budget = n_layers - 1 - c["lock_on_layer"] bw_px = round(budget / (n_layers - 1) * 430) draw.text((1492, ty), name, fill=color, font=mono_font(13)) - draw.rounded_rectangle((1492, ty + 20, 1492 + bw_px, ty + 32), radius=6, fill=color) - draw.text((1492 + bw_px + 10, ty + 17), f"{budget} layers · p {c['max_answer_p']:.2f}", fill=P["muted"], font=mono_font(12)) + draw.rounded_rectangle( + (1492, ty + 20, 1492 + bw_px, ty + 32), radius=6, fill=color + ) + draw.text( + (1492 + bw_px + 10, ty + 17), + f"{budget} layers · p {c['max_answer_p']:.2f}", + fill=P["muted"], + font=mono_font(12), + ) ty += 44 # ---- Rule plate (bottom right). plate = (1460, 740, 2136, 1000) - draw.rounded_rectangle(plate, radius=22, fill=P["panel"], outline=(255, 196, 68), width=2) + draw.rounded_rectangle( + plate, radius=22, fill=P["panel"], outline=(255, 196, 68), width=2 + ) draw.text((1492, 762), "the rule", fill=P["amber"], font=ui_font(24, True)) rule_lines = [ "lock_on(patch) = min L such that", @@ -273,10 +439,19 @@ def main() -> None: ] ry = 806 for line in rule_lines: - draw.text((1492, ry), line, fill=P["ink"] if line else P["muted"], font=mono_font(17)) + draw.text( + (1492, ry), line, fill=P["ink"] if line else P["muted"], font=mono_font(17) + ) ry += 32 - draw.text((1492, 1014), f"question: {q['q'][:60]}…", fill=P["muted"], font=ui_font(15)) - draw.text((1492, 1040), f"answer: “{q['answer_text']}” · condition: {args.condition} · generation: “{cond['generation']}”", fill=P["muted"], font=ui_font(15)) + draw.text( + (1492, 1014), f"question: {q['q'][:60]}…", fill=P["muted"], font=ui_font(15) + ) + draw.text( + (1492, 1040), + f"answer: “{q['answer_text']}” · condition: {args.condition} · generation: “{cond['generation']}”", + fill=P["muted"], + font=ui_font(15), + ) out = Path(args.out) out.parent.mkdir(parents=True, exist_ok=True) diff --git a/packages/snapcompact/research/snapcompact_logit_lens_dump.py b/packages/snapcompact/research/snapcompact_logit_lens_dump.py index 095c5de17..7953038aa 100644 --- a/packages/snapcompact/research/snapcompact_logit_lens_dump.py +++ b/packages/snapcompact/research/snapcompact_logit_lens_dump.py @@ -43,7 +43,11 @@ def main() -> None: args = ap.parse_args() import torch - from transformers import AutoProcessor, AutoTokenizer, Qwen2_5_VLForConditionalGeneration + from transformers import ( + AutoProcessor, + AutoTokenizer, + Qwen2_5_VLForConditionalGeneration, + ) out_dir = HERE / "results" / args.out img_dir = out_dir / "images" @@ -61,21 +65,55 @@ def main() -> None: img.save(img_dir / "image-carrier.png") print(f"loading {args.model_dir}", flush=True) - processor = AutoProcessor.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False) - tokenizer = AutoTokenizer.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True) - model = Qwen2_5_VLForConditionalGeneration.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, dtype=torch.bfloat16, device_map="auto").eval() + processor = AutoProcessor.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False + ) + tokenizer = AutoTokenizer.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True + ) + model = Qwen2_5_VLForConditionalGeneration.from_pretrained( + args.model_dir, + local_files_only=True, + trust_remote_code=True, + dtype=torch.bfloat16, + device_map="auto", + ).eval() device = next(model.parameters()).device - prompt = load_prompt("qa-image.md").format(cols=cols, rows=rows) + f"\n\nQuestion: {q['q']}\nAnswer with only the shortest extractive answer." - messages = [{"role": "user", "content": [{"type": "image", "image": img}, {"type": "text", "text": prompt}]}] - templated = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) + prompt = ( + load_prompt("qa-image.md").format(cols=cols, rows=rows) + + f"\n\nQuestion: {q['q']}\nAnswer with only the shortest extractive answer." + ) + messages = [ + { + "role": "user", + "content": [ + {"type": "image", "image": img}, + {"type": "text", "text": prompt}, + ], + } + ] + templated = processor.apply_chat_template( + messages, tokenize=False, add_generation_prompt=True + ) batch = processor(images=img, text=templated, return_tensors="pt") image_token_id = processor.tokenizer.convert_tokens_to_ids(processor.image_token) ids = batch["input_ids"][0].tolist() - image_positions = [i for i, token_id in enumerate(ids) if token_id == image_token_id] + image_positions = [ + i for i, token_id in enumerate(ids) if token_id == image_token_id + ] n_tokens = len(image_positions) grid = int(round(n_tokens**0.5)) - answer_indices = image_answer_token_indices(q["answer_start"], q["answer_end"], cols, cfg.adv, cfg.pitch, img.width, img.height, n_tokens) + answer_indices = image_answer_token_indices( + q["answer_start"], + q["answer_end"], + cols, + cfg.adv, + cfg.pitch, + img.width, + img.height, + n_tokens, + ) # Controls: blank-region tokens far from any text row boundary effects. control_indices = [] @@ -83,7 +121,9 @@ def main() -> None: row_far = (answer_indices[0] // grid + grid // 2) % grid for k in range(args.control_tokens): control_indices.append(row_far * grid + (answer_indices[0] % grid + k)) - track = [("answer", idx) for idx in answer_indices] + [("control", idx) for idx in control_indices] + track = [("answer", idx) for idx in answer_indices] + [ + ("control", idx) for idx in control_indices + ] track_positions = [image_positions[idx] for _kind, idx in track] batch = {k: (v.to(device) if hasattr(v, "to") else v) for k, v in batch.items()} @@ -92,7 +132,9 @@ def main() -> None: norm = model.model.language_model.norm lm_head = model.lm_head - answer_token_ids = tokenizer(q["answer_text"], add_special_tokens=False)["input_ids"] + answer_token_ids = tokenizer(q["answer_text"], add_special_tokens=False)[ + "input_ids" + ] answer_token_strs = [tokenizer.decode([t]) for t in answer_token_ids] lens: list[dict[str, Any]] = [] @@ -109,18 +151,34 @@ def main() -> None: "token_index": int(idx), "grid_rc": [int(idx // grid), int(idx % grid)], "top": [ - {"str": tokenizer.decode([int(topi[ti, k])]), "id": int(topi[ti, k]), "p": round(float(topv[ti, k]), 5)} + { + "str": tokenizer.decode([int(topi[ti, k])]), + "id": int(topi[ti, k]), + "p": round(float(topv[ti, k]), 5), + } for k in range(args.topk) ], - "answer_token_p": [round(float(probs[ti, t]), 6) for t in answer_token_ids], + "answer_token_p": [ + round(float(probs[ti, t]), 6) for t in answer_token_ids + ], } lens.append(entry) print(f"layer {layer} done", flush=True) dump = { "args": vars(args), - "question": {"q": q["q"], "answer_text": q["answer_text"], "answer_start": q["answer_start"], "answer_end": q["answer_end"]}, - "geometry": {"cols": cols, "rows": rows, "image_w": img.width, "image_h": img.height}, + "question": { + "q": q["q"], + "answer_text": q["answer_text"], + "answer_start": q["answer_start"], + "answer_end": q["answer_end"], + }, + "geometry": { + "cols": cols, + "rows": rows, + "image_w": img.width, + "image_h": img.height, + }, "image_tokens": n_tokens, "image_grid": grid, "token_pixel_size": 28, @@ -134,8 +192,20 @@ def main() -> None: (out_dir / "logit_lens.json").write_text(json.dumps(dump, indent=1)) # Quick console summary: best layer per answer token. for kind, idx in track: - best = max((e for e in lens if e["token_index"] == idx), key=lambda e: max(e["answer_token_p"])) - print(kind, idx, "best layer", best["layer"], "p", max(best["answer_token_p"]), "top1", best["top"][0]["str"]) + best = max( + (e for e in lens if e["token_index"] == idx), + key=lambda e: max(e["answer_token_p"]), + ) + print( + kind, + idx, + "best layer", + best["layer"], + "p", + max(best["answer_token_p"]), + "top1", + best["top"][0]["str"], + ) print(f"results -> {out_dir}") diff --git a/packages/snapcompact/research/snapcompact_logit_lens_viz.py b/packages/snapcompact/research/snapcompact_logit_lens_viz.py index 9c9b4eb15..046607bd6 100644 --- a/packages/snapcompact/research/snapcompact_logit_lens_viz.py +++ b/packages/snapcompact/research/snapcompact_logit_lens_viz.py @@ -31,8 +31,12 @@ PALETTE = { def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: for path in [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ]: if Path(path).exists(): return ImageFont.truetype(path, size) @@ -40,7 +44,10 @@ def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: def mono_font(size: int) -> ImageFont.ImageFont: - for path in ["/System/Library/Fonts/Monaco.ttf", "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf"]: + for path in [ + "/System/Library/Fonts/Monaco.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf", + ]: if Path(path).exists(): return ImageFont.truetype(path, size) return ImageFont.load_default() @@ -56,8 +63,13 @@ def heat_fill(p: float, hit: bool) -> tuple[int, int, int]: def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--result-dir", default=str(HERE / "results" / "qwen-logit-lens-q3")) - ap.add_argument("--out", default=str(HERE / "results" / "qwen-logit-lens-q3" / "logit-lens-grid.png")) + ap.add_argument( + "--result-dir", default=str(HERE / "results" / "qwen-logit-lens-q3") + ) + ap.add_argument( + "--out", + default=str(HERE / "results" / "qwen-logit-lens-q3" / "logit-lens-grid.png"), + ) ap.add_argument("--layer-step", type=int, default=1) args = ap.parse_args() result_dir = Path(args.result_dir) @@ -98,11 +110,23 @@ def main() -> None: gd = ImageDraw.Draw(glow) gd.ellipse((-240, -200, 900, 700), fill=(75, 220, 255, 26)) gd.ellipse((w - 1000, h - 800, w + 240, h + 200), fill=(255, 112, 72, 24)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(86))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(86)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) - draw.text((margin, 42), "QWEN LOGIT LENS — PIXELS BECOMING WORDS", fill=PALETTE["amber"], font=ui_font(24, True)) - draw.text((margin, 84), "Watch each patch decode into vocabulary", fill=PALETTE["ink"], font=ui_font(58, True)) + draw.text( + (margin, 42), + "QWEN LOGIT LENS — PIXELS BECOMING WORDS", + fill=PALETTE["amber"], + font=ui_font(24, True), + ) + draw.text( + (margin, 84), + "Watch each patch decode into vocabulary", + fill=PALETTE["ink"], + font=ui_font(58, True), + ) draw.text( (margin + 2, 156), f"Each column is one 28×28px visual token; each row is a decoder layer projected through the LM head. Green cells decode to a BPE piece of “{answer}”.", @@ -116,22 +140,45 @@ def main() -> None: patch_size = 108 for ci, idx in enumerate(track_indices): r, c = idx // grid, idx % grid - cell = resized.crop((c * px, r * px, (c + 1) * px, (r + 1) * px)).resize((patch_size, patch_size), Image.Resampling.NEAREST) + cell = resized.crop((c * px, r * px, (c + 1) * px, (r + 1) * px)).resize( + (patch_size, patch_size), Image.Resampling.NEAREST + ) cx = gx0 + ci * cell_w + (cell_w - patch_size) // 2 is_control = idx in dump["control_indices"] color = PALETTE["muted"] if is_control else PALETTE["orange"] - draw.rounded_rectangle((cx - 4, title_h + 26, cx + patch_size + 4, title_h + 34 + patch_size), radius=8, fill=(244, 242, 230), outline=color, width=3) + draw.rounded_rectangle( + (cx - 4, title_h + 26, cx + patch_size + 4, title_h + 34 + patch_size), + radius=8, + fill=(244, 242, 230), + outline=color, + width=3, + ) canvas.paste(cell, (cx, title_h + 30)) label = "control" if is_control else f"tok[{idx}]" tw = draw.textlength(label, font=mono_font(13)) - draw.text((cx + (patch_size - tw) / 2, title_h + 42 + patch_size), label, fill=color, font=mono_font(13)) - draw.text((margin, title_h + 30 + patch_size // 2 - 10), "input\npixels", fill=PALETTE["muted"], font=ui_font(15, True)) + draw.text( + (cx + (patch_size - tw) / 2, title_h + 42 + patch_size), + label, + fill=color, + font=mono_font(13), + ) + draw.text( + (margin, title_h + 30 + patch_size // 2 - 10), + "input\npixels", + fill=PALETTE["muted"], + font=ui_font(15, True), + ) # Grid rows. fnt = mono_font(14) for ri, layer in enumerate(layer_rows): y = gy0 + ri * cell_h - draw.text((margin + 24, y + 8), f"L{layer:02d}", fill=PALETTE["muted"], font=mono_font(13)) + draw.text( + (margin + 24, y + 8), + f"L{layer:02d}", + fill=PALETTE["muted"], + font=mono_font(13), + ) for ci, idx in enumerate(track_indices): e = by_token[idx][layer] top = e["top"][0] @@ -139,19 +186,47 @@ def main() -> None: p_ans = max(e["answer_token_p"]) x = gx0 + ci * cell_w fill = heat_fill(top["p"] if not hit else max(top["p"], p_ans), hit) - draw.rounded_rectangle((x + 2, y + 2, x + cell_w - 6, y + cell_h - 4), radius=6, fill=fill, outline=(32, 44, 53), width=1) + draw.rounded_rectangle( + (x + 2, y + 2, x + cell_w - 6, y + cell_h - 4), + radius=6, + fill=fill, + outline=(32, 44, 53), + width=1, + ) label = top["str"].replace("\n", "⏎").strip() or "·" if len(label) > 12: label = label[:11] + "…" - color = (220, 255, 190) if hit else PALETTE["ink"] if top["p"] > 0.05 else PALETTE["muted"] + color = ( + (220, 255, 190) + if hit + else PALETTE["ink"] + if top["p"] > 0.05 + else PALETTE["muted"] + ) draw.text((x + 10, y + 8), label, fill=color, font=fnt) if hit: - draw.text((x + cell_w - 52, y + 9), f"{p_ans:.2f}", fill=PALETTE["green"], font=mono_font(11)) + draw.text( + (x + cell_w - 52, y + 9), + f"{p_ans:.2f}", + fill=PALETTE["green"], + font=mono_font(11), + ) # Footer. fy = gy0 + len(layer_rows) * cell_h + 22 - draw.rounded_rectangle((margin, fy, w - margin, fy + 88), radius=18, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1) - draw.text((margin + 28, fy + 16), f"question: {q['q'][:88]}", fill=PALETTE["ink"], font=ui_font(19)) + draw.rounded_rectangle( + (margin, fy, w - margin, fy + 88), + radius=18, + fill=PALETTE["panel2"], + outline=(34, 48, 58), + width=1, + ) + draw.text( + (margin + 28, fy + 16), + f"question: {q['q'][:88]}", + fill=PALETTE["ink"], + font=ui_font(19), + ) draw.text( (margin + 28, fy + 50), f"gold answer “{answer}” = BPE {dump['answer_token_strs']} · logit lens = hidden state → final norm → LM head · {dump['image_tokens']:,} visual tokens total, showing the {len(dump['answer_indices'])} covering the answer + {len(dump['control_indices'])} blank-region controls", diff --git a/packages/snapcompact/research/snapcompact_materialize_sweep.py b/packages/snapcompact/research/snapcompact_materialize_sweep.py index 68ab4cfc5..de3e3bf31 100644 --- a/packages/snapcompact/research/snapcompact_materialize_sweep.py +++ b/packages/snapcompact/research/snapcompact_materialize_sweep.py @@ -39,12 +39,48 @@ class Condition: CONDITIONS = [ - Condition("base-8x13", FONTS["8x13"], "bw", 1, "baseline: glyphs straddle token cells on both axes"), - Condition("repeat2-color", FONTS["8x13"], "color", 2, "every line twice, consecutive rows in different hues"), - Condition("align-7x14", FontCfg("7x14a", "7x13", 7, 14), "bw", 1, "4 chars x 2 rows per token, no straddling"), - Condition("align-14x28", FontCfg("14x28a", "7x13", 14, 28, native=(7, 14)), "bw", 1, "2 chars x 1 row per token"), - Condition("align-28x28", FontCfg("28x28a", "8x13", 28, 28, native=(8, 13)), "bw", 1, "1 char per token"), - Condition("repeat2-align-14x28", FontCfg("14x28a", "7x13", 14, 28, native=(7, 14)), "color", 2, "aligned + repeated lines in hues"), + Condition( + "base-8x13", + FONTS["8x13"], + "bw", + 1, + "baseline: glyphs straddle token cells on both axes", + ), + Condition( + "repeat2-color", + FONTS["8x13"], + "color", + 2, + "every line twice, consecutive rows in different hues", + ), + Condition( + "align-7x14", + FontCfg("7x14a", "7x13", 7, 14), + "bw", + 1, + "4 chars x 2 rows per token, no straddling", + ), + Condition( + "align-14x28", + FontCfg("14x28a", "7x13", 14, 28, native=(7, 14)), + "bw", + 1, + "2 chars x 1 row per token", + ), + Condition( + "align-28x28", + FontCfg("28x28a", "8x13", 28, 28, native=(8, 13)), + "bw", + 1, + "1 char per token", + ), + Condition( + "repeat2-align-14x28", + FontCfg("14x28a", "7x13", 14, 28, native=(7, 14)), + "color", + 2, + "aligned + repeated lines in hues", + ), ] @@ -66,7 +102,16 @@ def build_layout(chunk: str, cols: int, rows: int, repeat: int) -> tuple[str, in return "".join(out), usable -def answer_token_indices(start: int, end: int, cols: int, adv: int, pitch: int, repeat: int, image_size: int, grid: int) -> list[int]: +def answer_token_indices( + start: int, + end: int, + cols: int, + adv: int, + pitch: int, + repeat: int, + image_size: int, + grid: int, +) -> list[int]: """Visual-token indices covering chars [start, end) under the layout.""" indices: set[int] = set() for i in range(start, end): @@ -98,7 +143,11 @@ def main() -> None: args = ap.parse_args() import torch - from transformers import AutoProcessor, AutoTokenizer, Qwen2_5_VLForConditionalGeneration + from transformers import ( + AutoProcessor, + AutoTokenizer, + Qwen2_5_VLForConditionalGeneration, + ) out_dir = HERE / "results" / args.out img_dir = out_dir / "images" @@ -112,17 +161,34 @@ def main() -> None: paras = squad.load_paragraphs(CACHE)[: args.limit_paras] flow, offsets = squad.build_flow(paras) base_chunk = flow[: min(len(flow), base_budget)] - questions = sample_answer_questions(paras, offsets, 0, len(base_chunk), 24, args.seed) + questions = sample_answer_questions( + paras, offsets, 0, len(base_chunk), 24, args.seed + ) q = questions[min(args.question_index, len(questions) - 1)] - print(f"question: {q['q']!r} answer: {q['answer_text']!r} @ {q['answer_start']}", flush=True) + print( + f"question: {q['q']!r} answer: {q['answer_text']!r} @ {q['answer_start']}", + flush=True, + ) print(f"loading {args.model_dir}", flush=True) - processor = AutoProcessor.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False) - tokenizer = AutoTokenizer.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True) - model = Qwen2_5_VLForConditionalGeneration.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, dtype=torch.bfloat16, device_map="auto").eval() + processor = AutoProcessor.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False + ) + tokenizer = AutoTokenizer.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True + ) + model = Qwen2_5_VLForConditionalGeneration.from_pretrained( + args.model_dir, + local_files_only=True, + trust_remote_code=True, + dtype=torch.bfloat16, + device_map="auto", + ).eval() device = next(model.parameters()).device image_token_id = processor.tokenizer.convert_tokens_to_ids(processor.image_token) - answer_token_ids = tokenizer(q["answer_text"], add_special_tokens=False)["input_ids"] + answer_token_ids = tokenizer(q["answer_text"], add_special_tokens=False)[ + "input_ids" + ] answer_id_set = set(answer_token_ids) answer_token_strs = [tokenizer.decode([t]) for t in answer_token_ids] norm = model.model.language_model.norm @@ -138,21 +204,47 @@ def main() -> None: img = render(render_text, cond.cfg, CACHE, args.size, cond.variant) img.save(img_dir / f"{cond.name}.png") - prompt = load_prompt("qa-image.md").format(cols=cols, rows=rows) + f"\n\nQuestion: {q['q']}\nAnswer with only the shortest extractive answer." - messages = [{"role": "user", "content": [{"type": "image", "image": img}, {"type": "text", "text": prompt}]}] - templated = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) + prompt = ( + load_prompt("qa-image.md").format(cols=cols, rows=rows) + + f"\n\nQuestion: {q['q']}\nAnswer with only the shortest extractive answer." + ) + messages = [ + { + "role": "user", + "content": [ + {"type": "image", "image": img}, + {"type": "text", "text": prompt}, + ], + } + ] + templated = processor.apply_chat_template( + messages, tokenize=False, add_generation_prompt=True + ) batch = processor(images=img, text=templated, return_tensors="pt") ids = batch["input_ids"][0].tolist() - image_positions = [i for i, token_id in enumerate(ids) if token_id == image_token_id] + image_positions = [ + i for i, token_id in enumerate(ids) if token_id == image_token_id + ] grid = int(round(len(image_positions) ** 0.5)) - track = answer_token_indices(q["answer_start"], q["answer_end"], cols, cond.cfg.adv, cond.cfg.pitch, cond.repeat, args.size, grid) + track = answer_token_indices( + q["answer_start"], + q["answer_end"], + cols, + cond.cfg.adv, + cond.cfg.pitch, + cond.repeat, + args.size, + grid, + ) track_positions = [image_positions[idx] for idx in track] batch = {k: (v.to(device) if hasattr(v, "to") else v) for k, v in batch.items()} with torch.no_grad(): fwd = model(**batch, output_hidden_states=True, use_cache=False) generated = model.generate(**batch, max_new_tokens=16, do_sample=False) - answer_gen = processor.batch_decode(generated[:, batch["input_ids"].shape[1] :], skip_special_tokens=True)[0].strip() + answer_gen = processor.batch_decode( + generated[:, batch["input_ids"].shape[1] :], skip_special_tokens=True + )[0].strip() layers_data: list[dict[str, Any]] = [] lock_on_layer: int | None = None @@ -175,7 +267,11 @@ def main() -> None: "top1_hit": bool(top1_hit), "best_token_index": track[best_idx], "best_token_top": [ - {"str": tokenizer.decode([int(ti[k])]), "p": round(float(tv[k]), 5)} for k in range(args.topk) + { + "str": tokenizer.decode([int(ti[k])]), + "p": round(float(tv[k]), 5), + } + for k in range(args.topk) ], } ) @@ -213,7 +309,12 @@ def main() -> None: summary = { "args": vars(args), - "question": {"q": q["q"], "answer_text": q["answer_text"], "answer_start": q["answer_start"], "answer_end": q["answer_end"]}, + "question": { + "q": q["q"], + "answer_text": q["answer_text"], + "answer_start": q["answer_start"], + "answer_end": q["answer_end"], + }, "answer_token_ids": answer_token_ids, "answer_token_strs": answer_token_strs, "conditions": conditions_out, diff --git a/packages/snapcompact/research/snapcompact_materialize_viz.py b/packages/snapcompact/research/snapcompact_materialize_viz.py index 0e5414f82..afbb024bd 100644 --- a/packages/snapcompact/research/snapcompact_materialize_viz.py +++ b/packages/snapcompact/research/snapcompact_materialize_viz.py @@ -34,8 +34,12 @@ SERIES = [ def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: for path in [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ]: if Path(path).exists(): return ImageFont.truetype(path, size) @@ -43,13 +47,22 @@ def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: def mono_font(size: int) -> ImageFont.ImageFont: - for path in ["/System/Library/Fonts/Monaco.ttf", "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf"]: + for path in [ + "/System/Library/Fonts/Monaco.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf", + ]: if Path(path).exists(): return ImageFont.truetype(path, size) return ImageFont.load_default() -def crop_answer_region(img_path: Path, cond: dict[str, Any], answer_start: int, answer_end: int, image_size: int = 1568) -> Image.Image: +def crop_answer_region( + img_path: Path, + cond: dict[str, Any], + answer_start: int, + answer_end: int, + image_size: int = 1568, +) -> Image.Image: img = Image.open(img_path).convert("RGB") cols = cond["cols"] adv = cond["adv"] @@ -64,14 +77,30 @@ def crop_answer_region(img_path: Path, cond: dict[str, Any], answer_start: int, x1 = min(image_size, (c1 + 1) * adv + 10 * adv) crop = img.crop((x0, y0, x1, y1)) d = ImageDraw.Draw(crop) - d.rectangle((c0 * adv - x0 - 2, row * repeat * pitch - y0 - 1, (c1 + 1) * adv - x0 + 2, (row * repeat + repeat) * pitch - y0 + 1), outline=(255, 112, 72), width=3) + d.rectangle( + ( + c0 * adv - x0 - 2, + row * repeat * pitch - y0 - 1, + (c1 + 1) * adv - x0 + 2, + (row * repeat + repeat) * pitch - y0 + 1, + ), + outline=(255, 112, 72), + width=3, + ) return crop def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--result-dir", default=str(HERE / "results" / "qwen-materialize-sweep-q3")) - ap.add_argument("--out", default=str(HERE / "results" / "qwen-materialize-sweep-q3" / "materialize-sweep.png")) + ap.add_argument( + "--result-dir", default=str(HERE / "results" / "qwen-materialize-sweep-q3") + ) + ap.add_argument( + "--out", + default=str( + HERE / "results" / "qwen-materialize-sweep-q3" / "materialize-sweep.png" + ), + ) args = ap.parse_args() result_dir = Path(args.result_dir) summary = json.loads((result_dir / "summary.json").read_text()) @@ -87,11 +116,23 @@ def main() -> None: gd = ImageDraw.Draw(glow) gd.ellipse((-240, -200, 900, 700), fill=(75, 220, 255, 25)) gd.ellipse((1240, 540, 2460, 1480), fill=(255, 112, 72, 24)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(86))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(86)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) - draw.text((64, 42), "QWEN MATERIALIZATION SWEEP — CAN RENDERING MOVE THE LAYER?", fill=(255, 196, 68), font=ui_font(24, True)) - draw.text((64, 84), "The depth is the model's; the clarity is yours", fill=PALETTE["ink"], font=ui_font(58, True)) + draw.text( + (64, 42), + "QWEN MATERIALIZATION SWEEP — CAN RENDERING MOVE THE LAYER?", + fill=(255, 196, 68), + font=ui_font(24, True), + ) + draw.text( + (64, 84), + "The depth is the model's; the clarity is yours", + fill=PALETTE["ink"], + font=ui_font(58, True), + ) draw.text( (66, 158), "Six renderings of the same passage. Logit-lens p(answer BPE) at the answer patch, by layer.\n" @@ -102,14 +143,23 @@ def main() -> None: # Main curve panel. panel = (64, 226, 1380, 900) - draw.rounded_rectangle(panel, radius=26, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((96, 248), "p(answer BPE) at the best answer patch, per layer", fill=PALETTE["ink"], font=ui_font(26, True)) + draw.rounded_rectangle( + panel, radius=26, fill=PALETTE["panel"], outline=(35, 49, 59), width=1 + ) + draw.text( + (96, 248), + "p(answer BPE) at the best answer patch, per layer", + fill=PALETTE["ink"], + font=ui_font(26, True), + ) gx0, gy0, gx1, gy1 = 150, 320, 1330, 800 n_layers = len(conditions["base-8x13"]["layers"]) for i in range(6): yy = gy0 + (gy1 - gy0) * i / 5 draw.line((gx0, yy, gx1, yy), fill=PALETTE["grid"], width=1) - draw.text((96, yy - 9), f"{1.0 - i / 5:.1f}", fill=PALETTE["muted"], font=ui_font(14)) + draw.text( + (96, yy - 9), f"{1.0 - i / 5:.1f}", fill=PALETTE["muted"], font=ui_font(14) + ) for name, color in SERIES: cond = conditions.get(name) if not cond: @@ -122,12 +172,39 @@ def main() -> None: draw.line(pts, fill=color, width=5 if name != "base-8x13" else 4, joint="curve") if cond["lock_on_layer"] is not None: lx = gx0 + (gx1 - gx0) * cond["lock_on_layer"] / (n_layers - 1) - draw.ellipse((lx - 7, gy1 - (gy1 - gy0) * min(1.0, cond["layers"][cond["lock_on_layer"]]["best_answer_p"]) - 7, lx + 7, gy1 - (gy1 - gy0) * min(1.0, cond["layers"][cond["lock_on_layer"]]["best_answer_p"]) + 7), outline=color, width=3) + draw.ellipse( + ( + lx - 7, + gy1 + - (gy1 - gy0) + * min(1.0, cond["layers"][cond["lock_on_layer"]]["best_answer_p"]) + - 7, + lx + 7, + gy1 + - (gy1 - gy0) + * min(1.0, cond["layers"][cond["lock_on_layer"]]["best_answer_p"]) + + 7, + ), + outline=color, + width=3, + ) draw.text((gx0, gy1 + 16), "layer 0", fill=PALETTE["muted"], font=ui_font(15)) - draw.text((gx1 - 76, gy1 + 16), f"layer {n_layers - 1}", fill=PALETTE["muted"], font=ui_font(15)) - draw.text((gx0 + 320, gy1 + 16), "rings mark lock-on (top-1 becomes an answer BPE)", fill=PALETTE["muted"], font=ui_font(15)) + draw.text( + (gx1 - 76, gy1 + 16), + f"layer {n_layers - 1}", + fill=PALETTE["muted"], + font=ui_font(15), + ) + draw.text( + (gx0 + 320, gy1 + 16), + "rings mark lock-on (top-1 becomes an answer BPE)", + fill=PALETTE["muted"], + font=ui_font(15), + ) legend_box = (1420, 226, 2136, 900) - draw.rounded_rectangle(legend_box, radius=26, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) + draw.rounded_rectangle( + legend_box, radius=26, fill=PALETTE["panel"], outline=(35, 49, 59), width=1 + ) draw.text((1452, 248), "conditions", fill=PALETTE["ink"], font=ui_font(26, True)) ly = 304 for name, color in SERIES: @@ -136,30 +213,77 @@ def main() -> None: continue draw.rounded_rectangle((1452, ly, 1452 + 26, ly + 10), radius=4, fill=color) draw.text((1492, ly - 9), name, fill=PALETTE["ink"], font=ui_font(21, True)) - draw.text((1492, ly + 19), cond["note"], fill=PALETTE["muted"], font=ui_font(14)) - draw.text((1492, ly + 42), f"lock-on L{cond['lock_on_layer']} · peak p {cond['max_answer_p']:.2f} · {cond['chars_per_token']} chars/token", fill=color, font=mono_font(14)) + draw.text( + (1492, ly + 19), cond["note"], fill=PALETTE["muted"], font=ui_font(14) + ) + draw.text( + (1492, ly + 42), + f"lock-on L{cond['lock_on_layer']} · peak p {cond['max_answer_p']:.2f} · {cond['chars_per_token']} chars/token", + fill=color, + font=mono_font(14), + ) ly += 96 # Condition cards with real crops. card_y = 938 card_w = 660 - draw.text((64, card_y - 24), "what the model actually saw (answer region outlined)", fill=PALETTE["ink"], font=ui_font(22, True)) - positions = [(64, card_y + 10), (64 + card_w + 24, card_y + 10), (64 + 2 * (card_w + 24), card_y + 10)] + draw.text( + (64, card_y - 24), + "what the model actually saw (answer region outlined)", + fill=PALETTE["ink"], + font=ui_font(22, True), + ) + positions = [ + (64, card_y + 10), + (64 + card_w + 24, card_y + 10), + (64 + 2 * (card_w + 24), card_y + 10), + ] featured = ["base-8x13", "align-28x28", "repeat2-align-14x28"] for (cx, cy), name in zip(positions, featured): cond = conditions.get(name) if not cond: continue color = dict(SERIES)[name] - draw.rounded_rectangle((cx, cy, cx + card_w, cy + 350), radius=20, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1) + draw.rounded_rectangle( + (cx, cy, cx + card_w, cy + 350), + radius=20, + fill=PALETTE["panel2"], + outline=(34, 48, 58), + width=1, + ) draw.text((cx + 22, cy + 14), name, fill=color, font=ui_font(23, True)) - draw.text((cx + 22, cy + 46), cond["note"], fill=PALETTE["muted"], font=ui_font(15)) - crop = crop_answer_region(result_dir / "images" / f"{name}.png", cond, q["answer_start"], q["answer_end"]) + draw.text( + (cx + 22, cy + 46), cond["note"], fill=PALETTE["muted"], font=ui_font(15) + ) + crop = crop_answer_region( + result_dir / "images" / f"{name}.png", + cond, + q["answer_start"], + q["answer_end"], + ) scale = min((card_w - 44) / crop.width, 200 / crop.height) - crop_r = crop.resize((round(crop.width * scale), round(crop.height * scale)), Image.Resampling.NEAREST) - draw.rounded_rectangle((cx + 20, cy + 76, cx + card_w - 20, cy + 286), radius=12, fill=(244, 242, 230)) - canvas.paste(crop_r, (cx + 22 + (card_w - 44 - crop_r.width) // 2, cy + 78 + (206 - crop_r.height) // 2)) - draw.text((cx + 22, cy + 300), f"lock-on L{cond['lock_on_layer']} · peak p {cond['max_answer_p']:.2f} · {cond['chars_per_token']} chars/token · gen “{cond['generation']}”", fill=PALETTE["ink"], font=ui_font(16, True)) + crop_r = crop.resize( + (round(crop.width * scale), round(crop.height * scale)), + Image.Resampling.NEAREST, + ) + draw.rounded_rectangle( + (cx + 20, cy + 76, cx + card_w - 20, cy + 286), + radius=12, + fill=(244, 242, 230), + ) + canvas.paste( + crop_r, + ( + cx + 22 + (card_w - 44 - crop_r.width) // 2, + cy + 78 + (206 - crop_r.height) // 2, + ), + ) + draw.text( + (cx + 22, cy + 300), + f"lock-on L{cond['lock_on_layer']} · peak p {cond['max_answer_p']:.2f} · {cond['chars_per_token']} chars/token · gen “{cond['generation']}”", + fill=PALETTE["ink"], + font=ui_font(16, True), + ) out = Path(args.out) out.parent.mkdir(parents=True, exist_ok=True) diff --git a/packages/snapcompact/research/snapcompact_pricing_viz.py b/packages/snapcompact/research/snapcompact_pricing_viz.py index 6e113791f..d4ce1249e 100644 --- a/packages/snapcompact/research/snapcompact_pricing_viz.py +++ b/packages/snapcompact/research/snapcompact_pricing_viz.py @@ -35,8 +35,12 @@ CARRY = [ def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: for path in [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ]: if Path(path).exists(): return ImageFont.truetype(path, size) @@ -44,7 +48,10 @@ def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: def mono_font(size: int) -> ImageFont.ImageFont: - for path in ["/System/Library/Fonts/Monaco.ttf", "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf"]: + for path in [ + "/System/Library/Fonts/Monaco.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf", + ]: if Path(path).exists(): return ImageFont.truetype(path, size) return ImageFont.load_default() @@ -64,45 +71,123 @@ def main() -> None: gd = ImageDraw.Draw(glow) gd.ellipse((-260, -240, 800, 560), fill=(75, 220, 255, 26)) gd.ellipse((1400, 400, 2460, 1240), fill=(255, 196, 68, 26)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(88))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(88)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) draw.text((64, 40), "THE BILLING MATH", fill=P["amber"], font=ui_font(24, True)) - draw.text((64, 80), "A flat fee per canvas, no matter what's inside", fill=P["ink"], font=ui_font(56, True)) - draw.text((66, 152), "Anthropic bills images at width × height ÷ 750 tokens. Text tokens scale with content; image tokens scale with pixels. Dense fonts exploit the gap.", fill=P["muted"], font=ui_font(22)) + draw.text( + (64, 80), + "A flat fee per canvas, no matter what's inside", + fill=P["ink"], + font=ui_font(56, True), + ) + draw.text( + (66, 152), + "Anthropic bills images at width × height ÷ 750 tokens. Text tokens scale with content; image tokens scale with pixels. Dense fonts exploit the gap.", + fill=P["muted"], + font=ui_font(22), + ) # Formula card. card = (64, 224, 700, 420) - draw.rounded_rectangle(card, radius=22, fill=P["panel"], outline=(35, 49, 59), width=1) + draw.rounded_rectangle( + card, radius=22, fill=P["panel"], outline=(35, 49, 59), width=1 + ) draw.text((96, 246), "flat fee per canvas", fill=P["cyan"], font=ui_font(21, True)) - draw.text((96, 286), "1568 × 1568 → 3,279 tokens", fill=P["ink"], font=mono_font(24)) - draw.text((96, 326), "2576 × 2576 → 4,950 tokens", fill=P["ink"], font=mono_font(24)) - draw.text((96, 372), "(2576 is silently downscaled 0.75x — still the best $/char)", fill=P["muted"], font=ui_font(15)) + draw.text( + (96, 286), "1568 × 1568 → 3,279 tokens", fill=P["ink"], font=mono_font(24) + ) + draw.text( + (96, 326), "2576 × 2576 → 4,950 tokens", fill=P["ink"], font=mono_font(24) + ) + draw.text( + (96, 372), + "(2576 is silently downscaled 0.75x — still the best $/char)", + fill=P["muted"], + font=ui_font(15), + ) # Cache card. card = (64, 452, 700, 660) - draw.rounded_rectangle(card, radius=22, fill=P["panel"], outline=(35, 49, 59), width=1) + draw.rounded_rectangle( + card, radius=22, fill=P["panel"], outline=(35, 49, 59), width=1 + ) draw.text((96, 474), "with prompt caching", fill=P["green"], font=ui_font(21, True)) - draw.text((96, 514), "marginal re-ask ≈ 333 tokens/turn", fill=P["ink"], font=mono_font(22)) - draw.text((96, 554), "measured: 753 in · 3,330 cache-write", fill=P["muted"], font=mono_font(17)) - draw.text((96, 584), "16,650 cache-read over six calls", fill=P["muted"], font=mono_font(17)) - draw.text((96, 620), "$0.18 cached vs $0.33 uncached", fill=P["amber"], font=mono_font(18)) + draw.text( + (96, 514), + "marginal re-ask ≈ 333 tokens/turn", + fill=P["ink"], + font=mono_font(22), + ) + draw.text( + (96, 554), + "measured: 753 in · 3,330 cache-write", + fill=P["muted"], + font=mono_font(17), + ) + draw.text( + (96, 584), + "16,650 cache-read over six calls", + fill=P["muted"], + font=mono_font(17), + ) + draw.text( + (96, 620), "$0.18 cached vs $0.33 uncached", fill=P["amber"], font=mono_font(18) + ) # Fine-print card. card = (64, 692, 700, 920) - draw.rounded_rectangle(card, radius=22, fill=P["panel"], outline=(255, 112, 72), width=1) + draw.rounded_rectangle( + card, radius=22, fill=P["panel"], outline=(255, 112, 72), width=1 + ) draw.text((96, 714), "the decode tax", fill=P["orange"], font=ui_font(21, True)) - draw.text((96, 754), "Models reason their way through dense", fill=P["muted"], font=ui_font(18)) - draw.text((96, 782), "pixels: 5–10x more thinking tokens than", fill=P["muted"], font=ui_font(18)) - draw.text((96, 810), "text. Input savings are real; total cost", fill=P["muted"], font=ui_font(18)) - draw.text((96, 838), "depends on output pricing. Cache + re-ask", fill=P["muted"], font=ui_font(18)) - draw.text((96, 866), "is where it always wins.", fill=P["ink"], font=ui_font(18, True)) + draw.text( + (96, 754), + "Models reason their way through dense", + fill=P["muted"], + font=ui_font(18), + ) + draw.text( + (96, 782), + "pixels: 5–10x more thinking tokens than", + fill=P["muted"], + font=ui_font(18), + ) + draw.text( + (96, 810), + "text. Input savings are real; total cost", + fill=P["muted"], + font=ui_font(18), + ) + draw.text( + (96, 838), + "depends on output pricing. Cache + re-ask", + fill=P["muted"], + font=ui_font(18), + ) + draw.text( + (96, 866), "is where it always wins.", fill=P["ink"], font=ui_font(18, True) + ) # Carry bars. panel = (760, 224, 2136, 920) - draw.rounded_rectangle(panel, radius=26, fill=P["panel"], outline=(35, 49, 59), width=1) - draw.text((796, 250), "text-token equivalent carried vs image tokens billed", fill=P["ink"], font=ui_font(26, True)) - draw.text((796, 290), "same content, two meters — the orange bar is what you'd pay as text; the cyan bar is what the PNG bills", fill=P["muted"], font=ui_font(17)) + draw.rounded_rectangle( + panel, radius=26, fill=P["panel"], outline=(35, 49, 59), width=1 + ) + draw.text( + (796, 250), + "text-token equivalent carried vs image tokens billed", + fill=P["ink"], + font=ui_font(26, True), + ) + draw.text( + (796, 290), + "same content, two meters — the orange bar is what you'd pay as text; the cyan bar is what the PNG bills", + fill=P["muted"], + font=ui_font(17), + ) bx0, bx1 = 1100, 1860 max_tokens = 25000 y = 360 @@ -112,20 +197,46 @@ def main() -> None: tw = round((bx1 - bx0) * text_tokens / max_tokens) bw = round((bx1 - bx0) * billed / max_tokens) draw.rounded_rectangle((bx0, y, bx0 + tw, y + 26), radius=9, fill=P["orange"]) - draw.text((bx0 + tw + 12, y + 2), f"{text_tokens:,} as text", fill=P["orange"], font=mono_font(15)) - draw.rounded_rectangle((bx0, y + 34, bx0 + bw, y + 60), radius=9, fill=P["cyan"]) + draw.text( + (bx0 + tw + 12, y + 2), + f"{text_tokens:,} as text", + fill=P["orange"], + font=mono_font(15), + ) + draw.rounded_rectangle( + (bx0, y + 34, bx0 + bw, y + 60), radius=9, fill=P["cyan"] + ) ratio = text_tokens / billed - draw.text((bx0 + bw + 12, y + 36), f"{billed:,} billed · {ratio:.1f}x", fill=P["cyan"], font=mono_font(15)) + draw.text( + (bx0 + bw + 12, y + 36), + f"{billed:,} billed · {ratio:.1f}x", + fill=P["cyan"], + font=mono_font(15), + ) y += 130 # Cached marginal bar. - draw.text((796, y + 6), "any font · cached re-ask", fill=P["ink"], font=mono_font(17)) - draw.text((796, y + 32), "image as cached prefix block", fill=P["muted"], font=ui_font(13)) + draw.text( + (796, y + 6), "any font · cached re-ask", fill=P["ink"], font=mono_font(17) + ) + draw.text( + (796, y + 32), "image as cached prefix block", fill=P["muted"], font=ui_font(13) + ) bw = max(6, round((bx1 - bx0) * 333 / max_tokens)) draw.rounded_rectangle((bx0, y + 14, bx0 + bw, y + 40), radius=9, fill=P["green"]) - draw.text((bx0 + bw + 12, y + 16), "≈ 333 tokens/turn · 30x", fill=P["green"], font=mono_font(15)) + draw.text( + (bx0 + bw + 12, y + 16), + "≈ 333 tokens/turn · 30x", + fill=P["green"], + font=mono_font(15), + ) y += 110 draw.line((796, y, 2100, y), fill=P["grid"], width=1) - draw.text((796, y + 16), "10,000 tokens of text, carried by 3,279 image tokens, amortizing to ~333 — that's the whole pitch.", fill=P["amber"], font=ui_font(19, True)) + draw.text( + (796, y + 16), + "10,000 tokens of text, carried by 3,279 image tokens, amortizing to ~333 — that's the whole pitch.", + fill=P["amber"], + font=ui_font(19, True), + ) out = Path(args.out) out.parent.mkdir(parents=True, exist_ok=True) diff --git a/packages/snapcompact/research/snapcompact_qwen_control_intervention.py b/packages/snapcompact/research/snapcompact_qwen_control_intervention.py index 5c6f1376a..e5f00bda4 100644 --- a/packages/snapcompact/research/snapcompact_qwen_control_intervention.py +++ b/packages/snapcompact/research/snapcompact_qwen_control_intervention.py @@ -50,8 +50,12 @@ PALETTE = { def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: for path in [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ]: if Path(path).exists(): return ImageFont.truetype(path, size) @@ -59,7 +63,10 @@ def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: def mono_font(size: int) -> ImageFont.ImageFont: - for path in ["/System/Library/Fonts/Monaco.ttf", "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf"]: + for path in [ + "/System/Library/Fonts/Monaco.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf", + ]: if Path(path).exists(): return ImageFont.truetype(path, size) return ImageFont.load_default() @@ -91,22 +98,62 @@ def make_text_prompt(chunk: str, q: dict[str, Any]) -> str: def make_image_prompt(cols: int, rows: int, q: dict[str, Any]) -> str: - return load_prompt("qa-image.md").format(cols=cols, rows=rows) + f"\n\nQuestion: {q['q']}\nAnswer with only the shortest extractive answer." + return ( + load_prompt("qa-image.md").format(cols=cols, rows=rows) + + f"\n\nQuestion: {q['q']}\nAnswer with only the shortest extractive answer." + ) -def carrier_map(model: Any, processor: Any, img: Image.Image, chunk: str, q: dict[str, Any], cols: int, rows: int, device: Any) -> tuple[np.ndarray, np.ndarray, dict[str, Any]]: - text_layers, text_pos, _ = run_text(model, processor, make_text_prompt(chunk, q), chunk, q["answer_start"], q["answer_end"], device) - image_layers, image_positions, image_meta, _ = run_image(model, processor, img, make_image_prompt(cols, rows, q), device) +def carrier_map( + model: Any, + processor: Any, + img: Image.Image, + chunk: str, + q: dict[str, Any], + cols: int, + rows: int, + device: Any, +) -> tuple[np.ndarray, np.ndarray, dict[str, Any]]: + text_layers, text_pos, _ = run_text( + model, + processor, + make_text_prompt(chunk, q), + chunk, + q["answer_start"], + q["answer_end"], + device, + ) + image_layers, image_positions, image_meta, _ = run_image( + model, processor, img, make_image_prompt(cols, rows, q), device + ) image_count = len(image_positions) - answer_indices = image_answer_token_indices(q["answer_start"], q["answer_end"], cols, 8, 13, img.width, img.height, image_count) + answer_indices = image_answer_token_indices( + q["answer_start"], + q["answer_end"], + cols, + 8, + 13, + img.width, + img.height, + image_count, + ) sims = [] answer_cos = [] for text_h, image_h in zip(text_layers, image_layers): - text_ans = text_h[text_pos["answer_start"] : text_pos["answer_end"]].mean(axis=0) + text_ans = text_h[text_pos["answer_start"] : text_pos["answer_end"]].mean( + axis=0 + ) image_tokens = image_h[image_positions] image_ans = image_tokens[answer_indices] if answer_indices else image_tokens - sims.append(cosine(np.repeat(text_ans[None, :], image_tokens.shape[0], axis=0), image_tokens).astype(np.float32, copy=False)) - answer_cos.append(float(cosine(text_ans[None, :], image_ans.mean(axis=0, keepdims=True))[0])) + sims.append( + cosine( + np.repeat(text_ans[None, :], image_tokens.shape[0], axis=0), + image_tokens, + ).astype(np.float32, copy=False) + ) + answer_cos.append( + float(cosine(text_ans[None, :], image_ans.mean(axis=0, keepdims=True))[0]) + ) raw = np.stack(sims, axis=0) excess = raw - np.median(raw, axis=1, keepdims=True) norm, lo, hi = normalize_heat(excess) @@ -138,13 +185,40 @@ def generate_with_intervention( ) -> str: import torch - messages = [{"role": "user", "content": [{"type": "image", "image": img}, {"type": "text", "text": prompt}]}] - templated = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) + messages = [ + { + "role": "user", + "content": [ + {"type": "image", "image": img}, + {"type": "text", "text": prompt}, + ], + } + ] + templated = processor.apply_chat_template( + messages, tokenize=False, add_generation_prompt=True + ) batch = processor(images=img, text=templated, return_tensors="pt") image_token_id = processor.tokenizer.convert_tokens_to_ids(processor.image_token) - image_positions = [i for i, token_id in enumerate(batch["input_ids"][0].tolist()) if token_id == image_token_id] + image_positions = [ + i + for i, token_id in enumerate(batch["input_ids"][0].tolist()) + if token_id == image_token_id + ] rng = random.Random(seed) - random_indices = sorted(rng.sample([i for i in range(len(image_positions)) if i not in set(answer_indices)], len(answer_indices))) if answer_indices else [] + random_indices = ( + sorted( + rng.sample( + [ + i + for i in range(len(image_positions)) + if i not in set(answer_indices) + ], + len(answer_indices), + ) + ) + if answer_indices + else [] + ) target_indices = ( answer_indices if mode == "answer_mean_patch" @@ -159,6 +233,7 @@ def generate_with_intervention( handle = None if target_positions: + def hook(_module: Any, inputs: tuple[Any, ...]) -> tuple[Any, ...]: hidden = inputs[0] if hidden.ndim == 3 and hidden.shape[1] > max(target_positions): @@ -166,13 +241,17 @@ def generate_with_intervention( if mode == "all_image_zero": patched[:, target_positions, :] = 0 else: - source_positions = [p for p in image_positions if p not in target_positions] + source_positions = [ + p for p in image_positions if p not in target_positions + ] mean_vec = hidden[:, source_positions, :].mean(dim=1, keepdim=True) patched[:, target_positions, :] = mean_vec return (patched, *inputs[1:]) return inputs - handle = model.model.language_model.layers[layer].register_forward_pre_hook(hook) + handle = model.model.language_model.layers[layer].register_forward_pre_hook( + hook + ) try: with torch.no_grad(): generated = model.generate(**batch, max_new_tokens=24, do_sample=False) @@ -183,7 +262,9 @@ def generate_with_intervention( return processor.batch_decode(new_tokens, skip_special_tokens=True)[0].strip() -def crop_answer(img: Image.Image, q: dict[str, Any], cols: int, adv: int = 8, pitch: int = 13) -> Image.Image: +def crop_answer( + img: Image.Image, q: dict[str, Any], cols: int, adv: int = 8, pitch: int = 13 +) -> Image.Image: start = q["answer_start"] end = q["answer_end"] row0 = max(0, start // cols - 5) @@ -196,20 +277,40 @@ def crop_answer(img: Image.Image, q: dict[str, Any], cols: int, adv: int = 8, pi bx1 = min(crop.width - 1, ((end - 1) % cols - col0 + 2) * adv) by0 = max(0, (start // cols - row0) * pitch - 1) by1 = min(crop.height - 1, ((end - 1) // cols - row0 + 1) * pitch + 1) - d.rounded_rectangle((bx0, by0, bx1, by1), radius=3, outline=PALETTE["orange"], width=3) + d.rounded_rectangle( + (bx0, by0, bx1, by1), radius=3, outline=PALETTE["orange"], width=3 + ) return crop -def paste_fit(canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int]) -> None: +def paste_fit( + canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int] +) -> None: x0, y0, x1, y1 = box scale = min((x1 - x0) / img.width, (y1 - y0) / img.height) - resized = img.resize((max(1, round(img.width * scale)), max(1, round(img.height * scale))), Image.Resampling.NEAREST) - canvas.paste(resized, (x0 + (x1 - x0 - resized.width) // 2, y0 + (y1 - y0 - resized.height) // 2)) + resized = img.resize( + (max(1, round(img.width * scale)), max(1, round(img.height * scale))), + Image.Resampling.NEAREST, + ) + canvas.paste( + resized, + (x0 + (x1 - x0 - resized.width) // 2, y0 + (y1 - y0 - resized.height) // 2), + ) -def draw_grid(draw: ImageDraw.ImageDraw, grid_values: np.ndarray, answer_indices: list[int], box: tuple[int, int, int, int], title: str, subtitle: str, color: tuple[int, int, int]) -> None: +def draw_grid( + draw: ImageDraw.ImageDraw, + grid_values: np.ndarray, + answer_indices: list[int], + box: tuple[int, int, int, int], + title: str, + subtitle: str, + color: tuple[int, int, int], +) -> None: x0, y0, x1, y1 = box - draw.rounded_rectangle(box, radius=22, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1) + draw.rounded_rectangle( + box, radius=22, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1 + ) draw.text((x0 + 20, y0 + 18), title, fill=color, font=ui_font(25, True)) draw.text((x0 + 20, y0 + 50), subtitle, fill=PALETTE["muted"], font=ui_font(15)) gx0, gy0, gx1, gy1 = x0 + 30, y0 + 84, x1 - 30, y1 - 28 @@ -229,10 +330,23 @@ def draw_grid(draw: ImageDraw.ImageDraw, grid_values: np.ndarray, answer_indices xb = round(gx0 + (c + 1) * cw) ya = round(gy0 + r * ch) yb = round(gy0 + (r + 1) * ch) - draw.rectangle((xa - 2, ya - 2, xb + 2, yb + 2), outline=PALETTE["orange"], width=2) + draw.rectangle( + (xa - 2, ya - 2, xb + 2, yb + 2), outline=PALETTE["orange"], width=2 + ) -def render_figure(out_path: Path, img: Image.Image, primary: dict[str, Any], distractor: dict[str, Any], primary_norm: np.ndarray, distractor_norm: np.ndarray, primary_meta: dict[str, Any], distractor_meta: dict[str, Any], generations: dict[str, str], cols: int) -> None: +def render_figure( + out_path: Path, + img: Image.Image, + primary: dict[str, Any], + distractor: dict[str, Any], + primary_norm: np.ndarray, + distractor_norm: np.ndarray, + primary_meta: dict[str, Any], + distractor_meta: dict[str, Any], + generations: dict[str, str], + cols: int, +) -> None: w, h = 2200, 1320 canvas = Image.new("RGB", (w, h), PALETTE["bg"]) draw = ImageDraw.Draw(canvas) @@ -242,34 +356,122 @@ def render_figure(out_path: Path, img: Image.Image, primary: dict[str, Any], dis gd = ImageDraw.Draw(glow) gd.ellipse((-260, -240, 900, 760), fill=(75, 220, 255, 28)) gd.ellipse((1240, 120, 2480, 1380), fill=(255, 112, 72, 27)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(86))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(86)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) - draw.text((64, 42), "QWEN SNAPCOMPACT CONTROL + INTERVENTION", fill=PALETTE["amber"], font=ui_font(24, True)) - draw.text((64, 84), "Ask a different thing; patch the hidden answer", fill=PALETTE["ink"], font=ui_font(62, True)) - draw.text((66, 166), "Same bitmap, two questions. Then patch answer-region image-token activations at the peak layer and watch generation change.", fill=PALETTE["muted"], font=ui_font(24)) + draw.text( + (64, 42), + "QWEN SNAPCOMPACT CONTROL + INTERVENTION", + fill=PALETTE["amber"], + font=ui_font(24, True), + ) + draw.text( + (64, 84), + "Ask a different thing; patch the hidden answer", + fill=PALETTE["ink"], + font=ui_font(62, True), + ) + draw.text( + (66, 166), + "Same bitmap, two questions. Then patch answer-region image-token activations at the peak layer and watch generation change.", + fill=PALETTE["muted"], + font=ui_font(24), + ) - draw.rounded_rectangle((64, 238, 616, 1234), radius=30, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((96, 270), "same image carrier", fill=PALETTE["ink"], font=ui_font(32, True)) - draw.text((96, 312), "Qwen2.5-VL-7B, 1568px bitmap", fill=PALETTE["muted"], font=ui_font(18)) - for label, q, y, color in [("PRIMARY", primary, 374, PALETTE["orange"]), ("DISTRACTOR", distractor, 658, PALETTE["cyan"] )]: + draw.rounded_rectangle( + (64, 238, 616, 1234), + radius=30, + fill=PALETTE["panel"], + outline=(35, 49, 59), + width=1, + ) + draw.text( + (96, 270), "same image carrier", fill=PALETTE["ink"], font=ui_font(32, True) + ) + draw.text( + (96, 312), + "Qwen2.5-VL-7B, 1568px bitmap", + fill=PALETTE["muted"], + font=ui_font(18), + ) + for label, q, y, color in [ + ("PRIMARY", primary, 374, PALETTE["orange"]), + ("DISTRACTOR", distractor, 658, PALETTE["cyan"]), + ]: draw.text((96, y), label, fill=color, font=ui_font(17, True)) crop = crop_answer(img, q, cols) - draw.rounded_rectangle((96, y + 34, 584, y + 194), radius=14, fill=(244, 242, 230), outline=color, width=3) + draw.rounded_rectangle( + (96, y + 34, 584, y + 194), + radius=14, + fill=(244, 242, 230), + outline=color, + width=3, + ) paste_fit(canvas, crop, (112, y + 48, 568, y + 180)) draw.text((96, y + 216), q["q"][:58], fill=PALETTE["ink"], font=ui_font(18)) - draw.text((96, y + 244), f"gold: {q['answer_text']}", fill=PALETTE["amber"], font=ui_font(22, True)) - draw.text((96, 1012), f"primary peak: L{primary_meta['peak_layer']} cosine {primary_meta['peak_cosine']:.3f}", fill=PALETTE["orange"], font=ui_font(20, True)) - draw.text((96, 1044), f"distractor peak: L{distractor_meta['peak_layer']} cosine {distractor_meta['peak_cosine']:.3f}", fill=PALETTE["cyan"], font=ui_font(20, True)) - draw.text((96, 1102), f"image tokens: {primary_meta['image_tokens']} ({primary_meta['image_grid']}×{primary_meta['image_grid']})", fill=PALETTE["muted"], font=ui_font(18)) + draw.text( + (96, y + 244), + f"gold: {q['answer_text']}", + fill=PALETTE["amber"], + font=ui_font(22, True), + ) + draw.text( + (96, 1012), + f"primary peak: L{primary_meta['peak_layer']} cosine {primary_meta['peak_cosine']:.3f}", + fill=PALETTE["orange"], + font=ui_font(20, True), + ) + draw.text( + (96, 1044), + f"distractor peak: L{distractor_meta['peak_layer']} cosine {distractor_meta['peak_cosine']:.3f}", + fill=PALETTE["cyan"], + font=ui_font(20, True), + ) + draw.text( + (96, 1102), + f"image tokens: {primary_meta['image_tokens']} ({primary_meta['image_grid']}×{primary_meta['image_grid']})", + fill=PALETTE["muted"], + font=ui_font(18), + ) grid = primary_meta["image_grid"] - draw_grid(draw, primary_norm[primary_meta["peak_layer"]].reshape(grid, grid), primary_meta["answer_indices"], (666, 238, 1386, 706), "primary question map", f"{primary['answer_text']} @ layer {primary_meta['peak_layer']} — orange box marks true answer", PALETTE["orange"]) - draw_grid(draw, distractor_norm[distractor_meta["peak_layer"]].reshape(grid, grid), distractor_meta["answer_indices"], (1420, 238, 2140, 706), "distractor question map", f"{distractor['answer_text']} @ layer {distractor_meta['peak_layer']} — map should move", PALETTE["cyan"]) + draw_grid( + draw, + primary_norm[primary_meta["peak_layer"]].reshape(grid, grid), + primary_meta["answer_indices"], + (666, 238, 1386, 706), + "primary question map", + f"{primary['answer_text']} @ layer {primary_meta['peak_layer']} — orange box marks true answer", + PALETTE["orange"], + ) + draw_grid( + draw, + distractor_norm[distractor_meta["peak_layer"]].reshape(grid, grid), + distractor_meta["answer_indices"], + (1420, 238, 2140, 706), + "distractor question map", + f"{distractor['answer_text']} @ layer {distractor_meta['peak_layer']} — map should move", + PALETTE["cyan"], + ) - draw.rounded_rectangle((666, 746, 2140, 1234), radius=30, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((704, 780), "activation patch test", fill=PALETTE["ink"], font=ui_font(34, True)) - draw.text((704, 822), "Before decoder layer 0, replace selected image-token residuals. Local answer patches test specificity; all-image zero is the sanity check.", fill=PALETTE["muted"], font=ui_font(20)) + draw.rounded_rectangle( + (666, 746, 2140, 1234), + radius=30, + fill=PALETTE["panel"], + outline=(35, 49, 59), + width=1, + ) + draw.text( + (704, 780), "activation patch test", fill=PALETTE["ink"], font=ui_font(34, True) + ) + draw.text( + (704, 822), + "Before decoder layer 0, replace selected image-token residuals. Local answer patches test specificity; all-image zero is the sanity check.", + fill=PALETTE["muted"], + font=ui_font(20), + ) rows = [ ("normal", generations["normal"], PALETTE["green"]), ("patch random region", generations["random_mean_patch"], PALETTE["cyan"]), @@ -278,9 +480,17 @@ def render_figure(out_path: Path, img: Image.Image, primary: dict[str, Any], dis ] y = 878 for label, text, color in rows: - draw.rounded_rectangle((704, y, 2078, y + 74), radius=18, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1) + draw.rounded_rectangle( + (704, y, 2078, y + 74), + radius=18, + fill=PALETTE["panel2"], + outline=(34, 48, 58), + width=1, + ) draw.text((730, y + 16), label.upper(), fill=color, font=ui_font(17, True)) - draw.text((1002, y + 15), text[:115], fill=PALETTE["ink"], font=ui_font(23, True)) + draw.text( + (1002, y + 15), text[:115], fill=PALETTE["ink"], font=ui_font(23, True) + ) y += 86 out_path.parent.mkdir(parents=True, exist_ok=True) @@ -314,7 +524,9 @@ def main() -> None: paras = squad.load_paragraphs(CACHE)[: args.limit_paras] flow, offsets = squad.build_flow(paras) chunk = flow[: min(len(flow), budget)] - questions = sample_answer_questions(paras, offsets, 0, len(chunk), args.qpc, args.seed) + questions = sample_answer_questions( + paras, offsets, 0, len(chunk), args.qpc, args.seed + ) if len(questions) < 2: raise SystemExit("not enough questions fit in chunk") primary = questions[min(args.question_index, len(questions) - 1)] @@ -325,21 +537,73 @@ def main() -> None: img.save(img_dir / "image-carrier.png") print(f"loading {args.model_dir}", flush=True) - processor = AutoProcessor.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False) - model = Qwen2_5_VLForConditionalGeneration.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, dtype=torch.bfloat16, device_map="auto").eval() + processor = AutoProcessor.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False + ) + model = Qwen2_5_VLForConditionalGeneration.from_pretrained( + args.model_dir, + local_files_only=True, + trust_remote_code=True, + dtype=torch.bfloat16, + device_map="auto", + ).eval() device = next(model.parameters()).device - primary_raw, primary_norm, primary_meta = carrier_map(model, processor, img, chunk, primary, cols, rows, device) - distractor_raw, distractor_norm, distractor_meta = carrier_map(model, processor, img, chunk, distractor, cols, rows, device) + primary_raw, primary_norm, primary_meta = carrier_map( + model, processor, img, chunk, primary, cols, rows, device + ) + distractor_raw, distractor_norm, distractor_meta = carrier_map( + model, processor, img, chunk, distractor, cols, rows, device + ) peak_layer = primary_meta["peak_layer"] prompt = make_image_prompt(cols, rows, primary) patch_layer = 0 generations = { - "normal": generate_with_intervention(model, processor, img, prompt, device, patch_layer, primary_meta["answer_indices"], "none", args.seed), - "random_mean_patch": generate_with_intervention(model, processor, img, prompt, device, patch_layer, primary_meta["answer_indices"], "random_mean_patch", args.seed), - "answer_mean_patch": generate_with_intervention(model, processor, img, prompt, device, patch_layer, primary_meta["answer_indices"], "answer_mean_patch", args.seed), - "all_image_zero": generate_with_intervention(model, processor, img, prompt, device, patch_layer, primary_meta["answer_indices"], "all_image_zero", args.seed), + "normal": generate_with_intervention( + model, + processor, + img, + prompt, + device, + patch_layer, + primary_meta["answer_indices"], + "none", + args.seed, + ), + "random_mean_patch": generate_with_intervention( + model, + processor, + img, + prompt, + device, + patch_layer, + primary_meta["answer_indices"], + "random_mean_patch", + args.seed, + ), + "answer_mean_patch": generate_with_intervention( + model, + processor, + img, + prompt, + device, + patch_layer, + primary_meta["answer_indices"], + "answer_mean_patch", + args.seed, + ), + "all_image_zero": generate_with_intervention( + model, + processor, + img, + prompt, + device, + patch_layer, + primary_meta["answer_indices"], + "all_image_zero", + args.seed, + ), } summary = { @@ -352,9 +616,26 @@ def main() -> None: "intervention_layer": patch_layer, "generations": generations, } - np.savez_compressed(out_dir / "control_intervention.npz", primary_raw=primary_raw, primary_norm=primary_norm, distractor_raw=distractor_raw, distractor_norm=distractor_norm) + np.savez_compressed( + out_dir / "control_intervention.npz", + primary_raw=primary_raw, + primary_norm=primary_norm, + distractor_raw=distractor_raw, + distractor_norm=distractor_norm, + ) (out_dir / "summary.json").write_text(json.dumps(summary, indent=1)) - render_figure(out_dir / "control-intervention.png", img, primary, distractor, primary_norm, distractor_norm, primary_meta, distractor_meta, generations, cols) + render_figure( + out_dir / "control-intervention.png", + img, + primary, + distractor, + primary_norm, + distractor_norm, + primary_meta, + distractor_meta, + generations, + cols, + ) print(json.dumps(summary, indent=1)) print(f"results -> {out_dir}") diff --git a/packages/snapcompact/research/snapcompact_qwen_spotlight_viz.py b/packages/snapcompact/research/snapcompact_qwen_spotlight_viz.py index 0e34ec882..36933ce28 100644 --- a/packages/snapcompact/research/snapcompact_qwen_spotlight_viz.py +++ b/packages/snapcompact/research/snapcompact_qwen_spotlight_viz.py @@ -33,8 +33,12 @@ PALETTE = { def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: for path in [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ]: if Path(path).exists(): return ImageFont.truetype(path, size) @@ -62,11 +66,15 @@ def normalize_positive(arr: np.ndarray) -> np.ndarray: def smooth_map(grid: np.ndarray, radius: float = 1.45) -> np.ndarray: g = normalize_positive(grid) - img = Image.fromarray(np.uint8(g * 255), mode="L").filter(ImageFilter.GaussianBlur(radius=radius)) + img = Image.fromarray(np.uint8(g * 255), mode="L").filter( + ImageFilter.GaussianBlur(radius=radius) + ) return np.asarray(img, dtype=np.float32) / 255.0 -def answer_bbox(indices: list[int], grid: int, image_w: int, image_h: int) -> tuple[int, int, int, int]: +def answer_bbox( + indices: list[int], grid: int, image_w: int, image_h: int +) -> tuple[int, int, int, int]: rows = [idx // grid for idx in indices] cols = [idx % grid for idx in indices] x0 = int(min(cols) / grid * image_w) @@ -76,11 +84,15 @@ def answer_bbox(indices: list[int], grid: int, image_w: int, image_h: int) -> tu return x0, y0, x1, y1 -def overlay_heat(base: Image.Image, heat: np.ndarray, bbox: tuple[int, int, int, int], theme: str) -> Image.Image: +def overlay_heat( + base: Image.Image, heat: np.ndarray, bbox: tuple[int, int, int, int], theme: str +) -> Image.Image: base_rgba = base.convert("RGBA") heat_img = Image.new("RGBA", base.size, (0, 0, 0, 0)) # Upscale smoothed 56x56 field to bitmap size; threshold softens static. - up = Image.fromarray(np.uint8(heat * 255), mode="L").resize(base.size, Image.Resampling.BICUBIC) + up = Image.fromarray(np.uint8(heat * 255), mode="L").resize( + base.size, Image.Resampling.BICUBIC + ) vals = np.asarray(up, dtype=np.float32) / 255.0 threshold = float(np.quantile(vals, 0.72)) vals = np.clip((vals - threshold) / max(1e-6, 1.0 - threshold), 0, 1) @@ -105,20 +117,41 @@ def overlay_heat(base: Image.Image, heat: np.ndarray, bbox: tuple[int, int, int, # spotlight ring around answer bbox x0, y0, x1, y1 = bbox pad = 24 - draw.rounded_rectangle((x0 - pad, y0 - pad, x1 + pad, y1 + pad), radius=16, outline=color, width=3) + draw.rounded_rectangle( + (x0 - pad, y0 - pad, x1 + pad, y1 + pad), radius=16, outline=color, width=3 + ) return out -def crop_box(img: Image.Image, bbox: tuple[int, int, int, int], pad: int = 180) -> Image.Image: +def crop_box( + img: Image.Image, bbox: tuple[int, int, int, int], pad: int = 180 +) -> Image.Image: x0, y0, x1, y1 = bbox - return img.crop((max(0, x0 - pad), max(0, y0 - pad), min(img.width, x1 + pad), min(img.height, y1 + pad))).convert("RGB") + return img.crop( + ( + max(0, x0 - pad), + max(0, y0 - pad), + min(img.width, x1 + pad), + min(img.height, y1 + pad), + ) + ).convert("RGB") -def paste_fit(canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int], resample: int = Image.Resampling.LANCZOS) -> None: +def paste_fit( + canvas: Image.Image, + img: Image.Image, + box: tuple[int, int, int, int], + resample: int = Image.Resampling.LANCZOS, +) -> None: x0, y0, x1, y1 = box scale = min((x1 - x0) / img.width, (y1 - y0) / img.height) - resized = img.resize((max(1, round(img.width * scale)), max(1, round(img.height * scale))), resample) - canvas.paste(resized, (x0 + (x1 - x0 - resized.width) // 2, y0 + (y1 - y0 - resized.height) // 2)) + resized = img.resize( + (max(1, round(img.width * scale)), max(1, round(img.height * scale))), resample + ) + canvas.paste( + resized, + (x0 + (x1 - x0 - resized.width) // 2, y0 + (y1 - y0 - resized.height) // 2), + ) def region_score(grid_map: np.ndarray, indices: list[int]) -> float: @@ -128,7 +161,9 @@ def region_score(grid_map: np.ndarray, indices: list[int]) -> float: return float(np.mean([flat[i] for i in indices if i < len(flat)])) -def random_region_scores(grid_map: np.ndarray, region_size: int, count: int = 600, seed: int = 7) -> np.ndarray: +def random_region_scores( + grid_map: np.ndarray, region_size: int, count: int = 600, seed: int = 7 +) -> np.ndarray: rng = random.Random(seed) flat = grid_map.ravel() scores = [] @@ -138,12 +173,26 @@ def random_region_scores(grid_map: np.ndarray, region_size: int, count: int = 60 return np.array(scores, dtype=np.float32) -def draw_score_card(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], title: str, score: float, random_scores: np.ndarray, color: tuple[int, int, int]) -> None: +def draw_score_card( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + title: str, + score: float, + random_scores: np.ndarray, + color: tuple[int, int, int], +) -> None: x0, y0, x1, y1 = box - draw.rounded_rectangle(box, radius=18, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1) + draw.rounded_rectangle( + box, radius=18, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1 + ) draw.text((x0 + 20, y0 + 18), title, fill=color, font=ui_font(22, True)) percentile = float((random_scores < score).mean() * 100) - draw.text((x0 + 20, y0 + 50), f"answer region beats {percentile:.0f}% of random same-size regions", fill=PALETTE["muted"], font=ui_font(16)) + draw.text( + (x0 + 20, y0 + 50), + f"answer region beats {percentile:.0f}% of random same-size regions", + fill=PALETTE["muted"], + font=ui_font(16), + ) gx0, gy0, gx1, gy1 = x0 + 28, y0 + 96, x1 - 28, y1 - 42 lo = float(min(random_scores.min(), score)) hi = float(max(random_scores.max(), score)) @@ -163,29 +212,68 @@ def draw_score_card(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], t draw.text((round(sx) - 34, gy0 - 36), "answer", fill=color, font=ui_font(15, True)) -def draw_generation_rows(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], generations: dict[str, str]) -> None: +def draw_generation_rows( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + generations: dict[str, str], +) -> None: x0, y0, x1, y1 = box - draw.rounded_rectangle(box, radius=24, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((x0 + 28, y0 + 24), "causal patch check", fill=PALETTE["ink"], font=ui_font(30, True)) - draw.text((x0 + 28, y0 + 60), "Patch before decoder layer 0; only the true answer-region patch changes the answer.", fill=PALETTE["muted"], font=ui_font(18)) + draw.rounded_rectangle( + box, radius=24, fill=PALETTE["panel"], outline=(35, 49, 59), width=1 + ) + draw.text( + (x0 + 28, y0 + 24), + "causal patch check", + fill=PALETTE["ink"], + font=ui_font(30, True), + ) + draw.text( + (x0 + 28, y0 + 60), + "Patch before decoder layer 0; only the true answer-region patch changes the answer.", + fill=PALETTE["muted"], + font=ui_font(18), + ) rows = [ ("normal", generations["normal"], PALETTE["green"]), ("random region patch", generations["random_mean_patch"], PALETTE["cyan"]), ("answer region patch", generations["answer_mean_patch"], PALETTE["red"]), - ("all image tokens zero", generations["all_image_zero"] or "∅", PALETTE["purple"]), + ( + "all image tokens zero", + generations["all_image_zero"] or "∅", + PALETTE["purple"], + ), ] y = y0 + 116 for label, text, color in rows: - draw.rounded_rectangle((x0 + 28, y, x1 - 28, y + 62), radius=14, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1) + draw.rounded_rectangle( + (x0 + 28, y, x1 - 28, y + 62), + radius=14, + fill=PALETTE["panel2"], + outline=(34, 48, 58), + width=1, + ) draw.text((x0 + 48, y + 17), label.upper(), fill=color, font=ui_font(15, True)) - draw.text((x0 + 330, y + 13), text[:70], fill=PALETTE["ink"], font=ui_font(24, True)) + draw.text( + (x0 + 330, y + 13), text[:70], fill=PALETTE["ink"], font=ui_font(24, True) + ) y += 78 def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--result-dir", default=str(HERE / "results" / "qwen-control-intervention-q3-d12-prehook")) - ap.add_argument("--out", default=str(HERE / "results" / "qwen-control-intervention-q3-d12-prehook" / "spotlight-control.png")) + ap.add_argument( + "--result-dir", + default=str(HERE / "results" / "qwen-control-intervention-q3-d12-prehook"), + ) + ap.add_argument( + "--out", + default=str( + HERE + / "results" + / "qwen-control-intervention-q3-d12-prehook" + / "spotlight-control.png" + ), + ) args = ap.parse_args() result_dir = Path(args.result_dir) @@ -207,15 +295,25 @@ def main() -> None: # for this question than for the other question? primary_contrast = smooth_map(primary_norm - distractor_norm) distractor_contrast = smooth_map(distractor_norm - primary_norm) - primary_bbox = answer_bbox(primary_meta["answer_indices"], grid, img.width, img.height) - distractor_bbox = answer_bbox(distractor_meta["answer_indices"], grid, img.width, img.height) + primary_bbox = answer_bbox( + primary_meta["answer_indices"], grid, img.width, img.height + ) + distractor_bbox = answer_bbox( + distractor_meta["answer_indices"], grid, img.width, img.height + ) primary_overlay = overlay_heat(img, primary_contrast, primary_bbox, "orange") distractor_overlay = overlay_heat(img, distractor_contrast, distractor_bbox, "cyan") primary_score = region_score(primary_contrast, primary_meta["answer_indices"]) - distractor_score = region_score(distractor_contrast, distractor_meta["answer_indices"]) - primary_random = random_region_scores(primary_contrast, len(primary_meta["answer_indices"]), seed=11) - distractor_random = random_region_scores(distractor_contrast, len(distractor_meta["answer_indices"]), seed=13) + distractor_score = region_score( + distractor_contrast, distractor_meta["answer_indices"] + ) + primary_random = random_region_scores( + primary_contrast, len(primary_meta["answer_indices"]), seed=11 + ) + distractor_random = random_region_scores( + distractor_contrast, len(distractor_meta["answer_indices"]), seed=13 + ) w, h = 2200, 1320 canvas = Image.new("RGB", (w, h), PALETTE["bg"]) @@ -226,34 +324,126 @@ def main() -> None: gd = ImageDraw.Draw(glow) gd.ellipse((-240, -200, 900, 700), fill=(255, 112, 72, 28)) gd.ellipse((1160, 120, 2460, 1340), fill=(75, 220, 255, 25)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(86))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(86)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) - draw.text((64, 42), "QWEN SNAPCOMPACT SPOTLIGHT", fill=PALETTE["amber"], font=ui_font(24, True)) - draw.text((64, 84), "Subtract the other question; the signal stops looking like static", fill=PALETTE["ink"], font=ui_font(58, True)) - draw.text((66, 160), "These are not raw activation carpets. Each overlay is prompt-specific excess: this question’s map minus the other question’s map, smoothed and thresholded.", fill=PALETTE["muted"], font=ui_font(23)) + draw.text( + (64, 42), + "QWEN SNAPCOMPACT SPOTLIGHT", + fill=PALETTE["amber"], + font=ui_font(24, True), + ) + draw.text( + (64, 84), + "Subtract the other question; the signal stops looking like static", + fill=PALETTE["ink"], + font=ui_font(58, True), + ) + draw.text( + (66, 160), + "These are not raw activation carpets. Each overlay is prompt-specific excess: this question’s map minus the other question’s map, smoothed and thresholded.", + fill=PALETTE["muted"], + font=ui_font(23), + ) # Overlay panels. - draw.rounded_rectangle((64, 230, 1068, 794), radius=28, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((96, 260), "primary prompt spotlight", fill=PALETTE["orange"], font=ui_font(31, True)) - draw.text((96, 298), f"{primary['q']} → {primary['answer_text']}", fill=PALETTE["muted"], font=ui_font(18)) + draw.rounded_rectangle( + (64, 230, 1068, 794), + radius=28, + fill=PALETTE["panel"], + outline=(35, 49, 59), + width=1, + ) + draw.text( + (96, 260), + "primary prompt spotlight", + fill=PALETTE["orange"], + font=ui_font(31, True), + ) + draw.text( + (96, 298), + f"{primary['q']} → {primary['answer_text']}", + fill=PALETTE["muted"], + font=ui_font(18), + ) primary_crop = crop_box(primary_overlay, primary_bbox, pad=300) paste_fit(canvas, primary_crop, (96, 342, 694, 760), Image.Resampling.LANCZOS) - draw.rounded_rectangle((720, 342, 1036, 760), radius=18, fill=(244, 242, 230), outline=PALETTE["orange"], width=3) - paste_fit(canvas, crop_box(primary_overlay, primary_bbox, pad=90), (736, 358, 1020, 744), Image.Resampling.LANCZOS) - draw.text((736, 724), "zoom: answer region", fill=PALETTE["orange"], font=ui_font(15, True)) + draw.rounded_rectangle( + (720, 342, 1036, 760), + radius=18, + fill=(244, 242, 230), + outline=PALETTE["orange"], + width=3, + ) + paste_fit( + canvas, + crop_box(primary_overlay, primary_bbox, pad=90), + (736, 358, 1020, 744), + Image.Resampling.LANCZOS, + ) + draw.text( + (736, 724), + "zoom: answer region", + fill=PALETTE["orange"], + font=ui_font(15, True), + ) - draw.rounded_rectangle((1132, 230, 2136, 794), radius=28, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((1164, 260), "distractor prompt spotlight", fill=PALETTE["cyan"], font=ui_font(31, True)) - draw.text((1164, 298), f"{distractor['q']} → {distractor['answer_text']}", fill=PALETTE["muted"], font=ui_font(18)) + draw.rounded_rectangle( + (1132, 230, 2136, 794), + radius=28, + fill=PALETTE["panel"], + outline=(35, 49, 59), + width=1, + ) + draw.text( + (1164, 260), + "distractor prompt spotlight", + fill=PALETTE["cyan"], + font=ui_font(31, True), + ) + draw.text( + (1164, 298), + f"{distractor['q']} → {distractor['answer_text']}", + fill=PALETTE["muted"], + font=ui_font(18), + ) distractor_crop = crop_box(distractor_overlay, distractor_bbox, pad=300) paste_fit(canvas, distractor_crop, (1164, 342, 1762, 760), Image.Resampling.LANCZOS) - draw.rounded_rectangle((1788, 342, 2104, 760), radius=18, fill=(244, 242, 230), outline=PALETTE["cyan"], width=3) - paste_fit(canvas, crop_box(distractor_overlay, distractor_bbox, pad=90), (1804, 358, 2088, 744), Image.Resampling.LANCZOS) - draw.text((1804, 724), "zoom: answer region", fill=PALETTE["cyan"], font=ui_font(15, True)) + draw.rounded_rectangle( + (1788, 342, 2104, 760), + radius=18, + fill=(244, 242, 230), + outline=PALETTE["cyan"], + width=3, + ) + paste_fit( + canvas, + crop_box(distractor_overlay, distractor_bbox, pad=90), + (1804, 358, 2088, 744), + Image.Resampling.LANCZOS, + ) + draw.text( + (1804, 724), "zoom: answer region", fill=PALETTE["cyan"], font=ui_font(15, True) + ) - draw_score_card(draw, (64, 836, 610, 1236), "primary answer-region score", primary_score, primary_random, PALETTE["orange"]) - draw_score_card(draw, (642, 836, 1188, 1236), "distractor answer-region score", distractor_score, distractor_random, PALETTE["cyan"]) + draw_score_card( + draw, + (64, 836, 610, 1236), + "primary answer-region score", + primary_score, + primary_random, + PALETTE["orange"], + ) + draw_score_card( + draw, + (642, 836, 1188, 1236), + "distractor answer-region score", + distractor_score, + distractor_random, + PALETTE["cyan"], + ) draw_generation_rows(draw, (1220, 836, 2136, 1236), summary["generations"]) # Save metrics alongside figure for the caption. @@ -261,7 +451,9 @@ def main() -> None: "primary_score": primary_score, "primary_percentile": float((primary_random < primary_score).mean() * 100), "distractor_score": distractor_score, - "distractor_percentile": float((distractor_random < distractor_score).mean() * 100), + "distractor_percentile": float( + (distractor_random < distractor_score).mean() * 100 + ), "primary_peak_layer": primary_layer, "distractor_peak_layer": distractor_layer, "generations": summary["generations"], diff --git a/packages/snapcompact/research/snapcompact_r2_chord.py b/packages/snapcompact/research/snapcompact_r2_chord.py index 0c687c034..fb24010ce 100755 --- a/packages/snapcompact/research/snapcompact_r2_chord.py +++ b/packages/snapcompact/research/snapcompact_r2_chord.py @@ -143,8 +143,12 @@ def center_bezier(a_deg, b_deg, pull=0.18, steps=60): p0, p3 = pol(a_deg, R_IN), pol(b_deg, R_IN) p1, p2 = p0 * pull, p3 * pull t = np.linspace(0, 1, steps)[:, None] - return ((1 - t) ** 3 * p0 + 3 * (1 - t) ** 2 * t * p1 - + 3 * (1 - t) * t ** 2 * p2 + t ** 3 * p3) + return ( + (1 - t) ** 3 * p0 + + 3 * (1 - t) ** 2 * t * p1 + + 3 * (1 - t) * t**2 * p2 + + t**3 * p3 + ) # ---------------------------------------------------------------- figure @@ -159,8 +163,9 @@ ax.axis("off") w_max = w.max() # mismatched ribbons first (thin, dim), then matched (amber, glowing) on top -order = sorted(((i, j) for i in range(n) for j in range(n)), - key=lambda ij: (ij[0] == ij[1], w[ij])) +order = sorted( + ((i, j) for i in range(n) for j in range(n)), key=lambda ij: (ij[0] == ij[1], w[ij]) +) for i, j in order: ls, rs = left_spans[i][j], right_spans[j][i] if ls is None or rs is None: @@ -175,57 +180,146 @@ for i, j in order: mid_r = 0.5 * (rs[0] + rs[1]) spine = center_bezier(mid_l, mid_r) for lw, al in ((26, 0.045), (14, 0.075), (7, 0.12)): - ax.plot(spine[:, 0], spine[:, 1], color=AMBER, lw=lw, alpha=al, - solid_capstyle="round", zorder=4) - ax.add_patch(PathPatch(path, facecolor=AMBER, edgecolor=AMBER, - lw=0.7, alpha=0.78, zorder=5)) + ax.plot( + spine[:, 0], + spine[:, 1], + color=AMBER, + lw=lw, + alpha=al, + solid_capstyle="round", + zorder=4, + ) + ax.add_patch( + PathPatch( + path, facecolor=AMBER, edgecolor=AMBER, lw=0.7, alpha=0.78, zorder=5 + ) + ) else: alpha = 0.10 + 0.45 * (val / w_max) - ax.add_patch(PathPatch(path, facecolor=MUTED, edgecolor="none", - alpha=alpha * 0.55, zorder=2)) + ax.add_patch( + PathPatch( + path, facecolor=MUTED, edgecolor="none", alpha=alpha * 0.55, zorder=2 + ) + ) # ---------------------------------------------------------------- node bands for i in range(n): - for centers, color, side in ((left_centers, CYAN, "L"), - (right_centers, ORANGE, "R")): + for centers, color, side in ( + (left_centers, CYAN, "L"), + (right_centers, ORANGE, "R"), + ): a0, a1 = seg_bounds(centers[i]) band = arc_points(a0, a1, R_OUT, 16) band_in = arc_points(a1, a0, R_IN, 16) poly = np.vstack([band, band_in]) - ax.add_patch(plt.Polygon(poly, closed=True, facecolor=color, - edgecolor="none", alpha=0.95, zorder=6)) + ax.add_patch( + plt.Polygon( + poly, + closed=True, + facecolor=color, + edgecolor="none", + alpha=0.95, + zorder=6, + ) + ) # ---------------------------------------------------------------- labels for i in range(n): txt = labels[i] - for centers, color, ha in ((left_centers, CYAN, "right"), - (right_centers, ORANGE, "left")): + for centers, color, ha in ( + (left_centers, CYAN, "right"), + (right_centers, ORANGE, "left"), + ): c = centers[i] p = pol(c, 1.05) - ax.text(p[0], p[1], txt, color=INK, fontsize=15.5, ha=ha, va="center", - zorder=8, family="DejaVu Sans") + ax.text( + p[0], + p[1], + txt, + color=INK, + fontsize=15.5, + ha=ha, + va="center", + zorder=8, + family="DejaVu Sans", + ) # small question index tick just inside the label - ax.text(p[0] + (0.018 if ha == "left" else -0.018), - p[1] - 0.052, f"Q{i + 1}", color=color, fontsize=10.5, - ha=ha, va="center", alpha=0.85, zorder=8) + ax.text( + p[0] + (0.018 if ha == "left" else -0.018), + p[1] - 0.052, + f"Q{i + 1}", + color=color, + fontsize=10.5, + ha=ha, + va="center", + alpha=0.85, + zorder=8, + ) # arc side headers -ax.text(*pol(180, 1.62), "TEXT CARRIER", color=CYAN, fontsize=21, - ha="center", va="center", rotation=90, weight="bold", alpha=0.95) -ax.text(*pol(180, 1.69), "5,219 prose tokens", color=MUTED, fontsize=13, - ha="center", va="center", rotation=90) -ax.text(*pol(0, 1.62), "IMAGE CARRIER", color=ORANGE, fontsize=21, - ha="center", va="center", rotation=-90, weight="bold", alpha=0.95) -ax.text(*pol(0, 1.69), "same passage, rendered as pixels", color=MUTED, - fontsize=13, ha="center", va="center", rotation=-90) +ax.text( + *pol(180, 1.62), + "TEXT CARRIER", + color=CYAN, + fontsize=21, + ha="center", + va="center", + rotation=90, + weight="bold", + alpha=0.95, +) +ax.text( + *pol(180, 1.69), + "5,219 prose tokens", + color=MUTED, + fontsize=13, + ha="center", + va="center", + rotation=90, +) +ax.text( + *pol(0, 1.62), + "IMAGE CARRIER", + color=ORANGE, + fontsize=21, + ha="center", + va="center", + rotation=-90, + weight="bold", + alpha=0.95, +) +ax.text( + *pol(0, 1.69), + "same passage, rendered as pixels", + color=MUTED, + fontsize=13, + ha="center", + va="center", + rotation=-90, +) # ---------------------------------------------------------------- titles -fig.text(0.5, 0.965, "ONE MEMORY, TWO CARRIERS", color=INK, fontsize=34, - ha="center", va="center", weight="bold", family="DejaVu Sans") -fig.text(0.5, 0.932, - "Cross-carrier cosine of answer states at layer 19 -- every text" - " question finds its image twin (Qwen2.5-VL-7B)", - color=MUTED, fontsize=16.5, ha="center", va="center") +fig.text( + 0.5, + 0.965, + "ONE MEMORY, TWO CARRIERS", + color=INK, + fontsize=34, + ha="center", + va="center", + weight="bold", + family="DejaVu Sans", +) +fig.text( + 0.5, + 0.932, + "Cross-carrier cosine of answer states at layer 19 -- every text" + " question finds its image twin (Qwen2.5-VL-7B)", + color=MUTED, + fontsize=16.5, + ha="center", + va="center", +) # ---------------------------------------------------------------- stat plate plate = fig.add_axes([0.035, 0.05, 0.215, 0.135]) @@ -236,8 +330,15 @@ for s in plate.spines.values(): s.set_color("#1c242e") plate.set_xticks([]) plate.set_yticks([]) -plate.text(0.5, 0.84, f"LAYER {LAYER} -- BEST SEPARATION", color=MUTED, - fontsize=12.5, ha="center", va="center") +plate.text( + 0.5, + 0.84, + f"LAYER {LAYER} -- BEST SEPARATION", + color=MUTED, + fontsize=12.5, + ha="center", + va="center", +) stats = ( (f"{matched_mean:+.2f}", "matched cosine", AMBER), (f"{mismatched_mean:+.2f}", "mismatched", MUTED), @@ -245,21 +346,29 @@ stats = ( ) for k, (val, lab, color) in enumerate(stats): x = 0.18 + 0.32 * k - plate.text(x, 0.48, val, color=color, fontsize=24, ha="center", - va="center", weight="bold") - plate.text(x, 0.18, lab, color=MUTED, fontsize=12, ha="center", - va="center") + plate.text( + x, 0.48, val, color=color, fontsize=24, ha="center", va="center", weight="bold" + ) + plate.text(x, 0.18, lab, color=MUTED, fontsize=12, ha="center", va="center") # footnote -fig.text(0.5, 0.012, - "Ribbon width/opacity = cosine(text_state_i, image_state_j)," - " negatives clipped; amber = matched pair (i = j)." - f" RSA r = {best['rsa_pearson']:.2f}.", - color=MUTED, fontsize=12.5, ha="center", va="center") +fig.text( + 0.5, + 0.012, + "Ribbon width/opacity = cosine(text_state_i, image_state_j)," + " negatives clipped; amber = matched pair (i = j)." + f" RSA r = {best['rsa_pearson']:.2f}.", + color=MUTED, + fontsize=12.5, + ha="center", + va="center", +) os.makedirs(OUT_DIR, exist_ok=True) out_path = os.path.join(OUT_DIR, "chord.png") fig.savefig(out_path, dpi=100, facecolor=BG) print(f"wrote {out_path}") -print(f"matched={matched_mean:.4f} mismatched={mismatched_mean:.4f} " - f"retrieval={retrieved}/{n}") +print( + f"matched={matched_mean:.4f} mismatched={mismatched_mean:.4f} " + f"retrieval={retrieved}/{n}" +) diff --git a/packages/snapcompact/research/snapcompact_r2_crystal.py b/packages/snapcompact/research/snapcompact_r2_crystal.py index fbe894bfc..fd0e417ce 100755 --- a/packages/snapcompact/research/snapcompact_r2_crystal.py +++ b/packages/snapcompact/research/snapcompact_r2_crystal.py @@ -155,13 +155,21 @@ def render_frame( badge = f"'{ANSWER_BPE}' \u00b7 p={p_ans:.2f} \u00b7 CRYSTALLIZED" bw = text_w(F_BADGE, badge) bx = W - bw - 76 - d.rounded_rectangle((bx, 22, bx + bw + 36, 56), radius=17, fill=(38, 30, 10), outline=AMBER, width=2) + d.rounded_rectangle( + (bx, 22, bx + bw + 36, 56), + radius=17, + fill=(38, 30, 10), + outline=AMBER, + width=2, + ) d.text((bx + 18, 29), badge, font=F_BADGE, fill=AMBER) # ---------------------------------------------------------- context strip sx, sy = (W - strip.width) // 2, 96 img.paste(strip, (sx, sy)) - d.rectangle((sx, sy, sx + strip.width - 1, sy + strip.height - 1), outline=PANEL_EDGE) + d.rectangle( + (sx, sy, sx + strip.width - 1, sy + strip.height - 1), outline=PANEL_EDGE + ) col = AMBER if locked else CYAN d.rectangle( (sx + strip_cell, sy, sx + strip_cell + strip.height, sy + strip.height - 1), @@ -169,7 +177,12 @@ def render_frame( width=3, ) cap = "carrier row 5 \u00b7 the model never sees glyphs \u2014 only these pixels" - d.text(((W - text_w(F_TINY, cap)) / 2, sy + strip.height + 6), cap, font=F_TINY, fill=DIM) + d.text( + ((W - text_w(F_TINY, cap)) / 2, sy + strip.height + 6), + cap, + font=F_TINY, + fill=DIM, + ) # ---------------------------------------------------------- main row top_y = 232 @@ -178,7 +191,12 @@ def render_frame( img.paste(patch.resize((ps, ps), Image.NEAREST), (px, py)) d.rectangle((px - 1, py - 1, px + ps, py + ps), outline=col, width=2) d.text((px, py + ps + 10), f"visual token #{TOKEN_IDX}", font=F_LABEL_B, fill=INK) - d.text((px, py + ps + 28), "28\u00d728 px \u00b7 reads: \u2018\"sp\u2019 / \u2018and\u2019", font=F_TINY, fill=MUTED) + d.text( + (px, py + ps + 28), + '28\u00d728 px \u00b7 reads: \u2018"sp\u2019 / \u2018and\u2019', + font=F_TINY, + fill=MUTED, + ) # layer counter ---------------------------------------------------- cx = 392 @@ -187,17 +205,34 @@ def render_frame( d.text((cx, top_y + 16), num, font=F_LAYER, fill=AMBER if locked else INK) d.text((cx + text_w(F_LAYER, num) + 8, top_y + 58), "/28", font=F_STAGE, fill=DIM) stage = stage_for(layer) - d.text((cx, top_y + 98), stage.upper(), font=F_STAGE, fill=GREEN if final else (AMBER if locked else MUTED)) + d.text( + (cx, top_y + 98), + stage.upper(), + font=F_STAGE, + fill=GREEN if final else (AMBER if locked else MUTED), + ) ry = top_y + 134 # mini rail of 29 ticks for i in range(29): tx = cx + i * 5 - d.rectangle((tx, ry, tx + 3, ry + 10), fill=AMBER if i <= layer else (40, 48, 58)) - d.text((cx, ry + 18), f"p('{ANSWER_BPE}') = {p_ans:.4f}", font=F_NUM, fill=AMBER if p_ans > 0.01 else DIM) + d.rectangle( + (tx, ry, tx + 3, ry + 10), fill=AMBER if i <= layer else (40, 48, 58) + ) + d.text( + (cx, ry + 18), + f"p('{ANSWER_BPE}') = {p_ans:.4f}", + font=F_NUM, + fill=AMBER if p_ans > 0.01 else DIM, + ) # top-5 panel ------------------------------------------------------ tx0, ty0, tx1, ty1 = 580, top_y - 12, 1160, top_y + 318 d.rounded_rectangle((tx0, ty0, tx1, ty1), radius=8, fill=PANEL, outline=PANEL_EDGE) - d.text((tx0 + 18, ty0 + 12), "TOP-5 DECODED VOCAB TOKENS \u00b7 what this patch \u201cmeans\u201d so far", font=F_LABEL, fill=MUTED) + d.text( + (tx0 + 18, ty0 + 12), + "TOP-5 DECODED VOCAB TOKENS \u00b7 what this patch \u201cmeans\u201d so far", + font=F_LABEL, + fill=MUTED, + ) bar_x = tx0 + 230 bar_max = tx1 - bar_x - 86 scale = 0.45 # fixed probability scale across all frames @@ -207,12 +242,24 @@ def render_frame( tok_s = sanitize(t["str"]) if len(tok_s) > 16: tok_s = tok_s[:15] + "\u2026" - d.text((tx0 + 18, yy), f"'{tok_s}'", font=F_TOK_B if is_ans else F_TOK, fill=AMBER if is_ans else INK) + d.text( + (tx0 + 18, yy), + f"'{tok_s}'", + font=F_TOK_B if is_ans else F_TOK, + fill=AMBER if is_ans else INK, + ) bw = max(2, int(min(t["p"] / scale, 1.0) * bar_max)) - d.rectangle((bar_x, yy + 4, bar_x + bw, yy + 18), fill=AMBER if is_ans else (58, 70, 84)) + d.rectangle( + (bar_x, yy + 4, bar_x + bw, yy + 18), fill=AMBER if is_ans else (58, 70, 84) + ) if is_ans and heat > 0.3: d.rectangle((bar_x, yy + 4, bar_x + bw, yy + 18), outline=INK) - d.text((bar_x + bw + 10, yy + 3), f"{t['p']:.3f}", font=F_NUM, fill=AMBER if is_ans else MUTED) + d.text( + (bar_x + bw + 10, yy + 3), + f"{t['p']:.3f}", + font=F_NUM, + fill=AMBER if is_ans else MUTED, + ) d.text((tx0 + 18, yy + 24), f"id {t['id']}", font=F_TINY, fill=DIM) # ---------------------------------------------------------- bottom row @@ -220,7 +267,12 @@ def render_frame( # confidence meter for 'acular' mx0, mx1 = 40, 730 d.rounded_rectangle((mx0, by0, mx1, by1), radius=8, fill=PANEL, outline=PANEL_EDGE) - d.text((mx0 + 16, by0 + 8), f"CONFIDENCE \u00b7 p('{ANSWER_BPE}') across layers", font=F_LABEL, fill=MUTED) + d.text( + (mx0 + 16, by0 + 8), + f"CONFIDENCE \u00b7 p('{ANSWER_BPE}') across layers", + font=F_LABEL, + fill=MUTED, + ) leg_x = mx1 - 130 d.rectangle((leg_x, by0 + 12, leg_x + 14, by0 + 15), fill=AMBER) d.text((leg_x + 20, by0 + 6), "answer", font=F_TINY, fill=AMBER) @@ -241,8 +293,12 @@ def render_frame( def ys(p: float) -> float: return ch_y1 - min(p, p_max) / p_max * (ch_y1 - ch_y0) - pts = [(xs(l), ys(target[l]["answer_token_p"][ANSWER_SLOT])) for l in range(layer + 1)] - cpts = [(xs(l), ys(control[l]["answer_token_p"][ANSWER_SLOT])) for l in range(layer + 1)] + pts = [ + (xs(l), ys(target[l]["answer_token_p"][ANSWER_SLOT])) for l in range(layer + 1) + ] + cpts = [ + (xs(l), ys(control[l]["answer_token_p"][ANSWER_SLOT])) for l in range(layer + 1) + ] if len(cpts) > 1: d.line(cpts, fill=(60, 70, 80), width=2) if len(pts) > 1: @@ -252,21 +308,36 @@ def render_frame( hx, hy = pts[-1] d.ellipse((hx - 5, hy - 5, hx + 5, hy + 5), fill=AMBER if p_ans > 0.01 else MUTED) head = f"{p_ans:.2f}" if p_ans >= 0.005 else f"{p_ans:.4f}" - d.text((min(hx + 8, ch_x1 - 8), hy - 18), head, font=F_NUM, fill=AMBER if p_ans > 0.01 else MUTED) + d.text( + (min(hx + 8, ch_x1 - 8), hy - 18), + head, + font=F_NUM, + fill=AMBER if p_ans > 0.01 else MUTED, + ) for ml in (24, 28): if layer >= ml: mlx = xs(ml) - d.line((mlx, ch_y1, mlx, ys(target[ml]["answer_token_p"][ANSWER_SLOT])), fill=(90, 72, 30)) + d.line( + (mlx, ch_y1, mlx, ys(target[ml]["answer_token_p"][ANSWER_SLOT])), + fill=(90, 72, 30), + ) d.text((mlx - 10, ch_y1 + 6), f"L{ml}", font=F_TINY, fill=AMBER) d.text((ch_x0, ch_y1 + 6), "L0", font=F_TINY, fill=DIM) # control panel ---------------------------------------------------- kx0, kx1 = 760, 1160 d.rounded_rectangle((kx0, by0, kx1, by1), radius=8, fill=PANEL, outline=PANEL_EDGE) - d.text((kx0 + 16, by0 + 8), "CONTROL \u00b7 token #" + str(ce["token_index"]), font=F_LABEL, fill=MUTED) + d.text( + (kx0 + 16, by0 + 8), + "CONTROL \u00b7 token #" + str(ce["token_index"]), + font=F_LABEL, + fill=MUTED, + ) cps = 60 img.paste(ctrl_patch.resize((cps, cps), Image.NEAREST), (kx0 + 16, by0 + 30)) - d.rectangle((kx0 + 15, by0 + 29, kx0 + 16 + cps, by0 + 30 + cps), outline=PANEL_EDGE) + d.rectangle( + (kx0 + 15, by0 + 29, kx0 + 16 + cps, by0 + 30 + cps), outline=PANEL_EDGE + ) ct = ce["top"][0] ct_s = sanitize(ct["str"]) if len(ct_s) > 12: @@ -275,9 +346,19 @@ def render_frame( d.text((lx, by0 + 30), "top-1: ", font=F_NUM, fill=INK) tx = lx + text_w(F_NUM, "top-1: ") d.text((tx, by0 + 30), f"'{ct_s}'", font=F_TOK_S, fill=INK) - d.text((tx + text_w(F_TOK_S, f"'{ct_s}'") + 10, by0 + 30), f"{ct['p']:.3f}", font=F_NUM, fill=INK) + d.text( + (tx + text_w(F_TOK_S, f"'{ct_s}'") + 10, by0 + 30), + f"{ct['p']:.3f}", + font=F_NUM, + fill=INK, + ) d.text((lx, by0 + 52), f"p('{ANSWER_BPE}') = {p_ctrl:.5f}", font=F_NUM, fill=MUTED) - d.text((lx, by0 + 74), "still noise \u2713" if p_ctrl < 0.01 else "?!", font=F_LABEL_B, fill=GREEN) + d.text( + (lx, by0 + 74), + "still noise \u2713" if p_ctrl < 0.01 else "?!", + font=F_LABEL_B, + fill=GREEN, + ) d.text((kx0 + 224, by0 + 8), "never converges to the answer", font=F_TINY, fill=DIM) # ---------------------------------------------------------- glow @@ -285,11 +366,23 @@ def render_frame( glow = Image.new("RGB", (W, H), (0, 0, 0)) gd = ImageDraw.Draw(glow) a = int(70 + 110 * heat) - gd.rectangle((px - 6, py - 6, px + ps + 5, py + ps + 5), outline=(a, int(a * 0.77), int(a * 0.27)), width=10) + gd.rectangle( + (px - 6, py - 6, px + ps + 5, py + ps + 5), + outline=(a, int(a * 0.77), int(a * 0.27)), + width=10, + ) if final: - gd.rectangle((px - 14, py - 14, px + ps + 13, py + ps + 13), outline=(a, int(a * 0.77), int(a * 0.27)), width=8) + gd.rectangle( + (px - 14, py - 14, px + ps + 13, py + ps + 13), + outline=(a, int(a * 0.77), int(a * 0.27)), + width=8, + ) glow = glow.filter(ImageFilter.GaussianBlur(12 if final else 8)) - img = Image.composite(Image.new("RGB", (W, H), AMBER), img, glow.convert("L").point(lambda v: min(v, 140))) + img = Image.composite( + Image.new("RGB", (W, H), AMBER), + img, + glow.convert("L").point(lambda v: min(v, 140)), + ) return img @@ -309,8 +402,12 @@ def main() -> None: frames, durations = [], [] for layer in range(29): - fr = render_frame(layer, data, target, control, patch, ctrl_patch, strip, strip_cell) - frames.append(fr.quantize(colors=256, method=Image.MEDIANCUT, dither=Image.Dither.NONE)) + fr = render_frame( + layer, data, target, control, patch, ctrl_patch, strip, strip_cell + ) + frames.append( + fr.quantize(colors=256, method=Image.MEDIANCUT, dither=Image.Dither.NONE) + ) if layer < 23: durations.append(220) elif layer < 28: @@ -329,7 +426,9 @@ def main() -> None: optimize=False, ) final_png = OUT_DIR / "crystal_final.png" - render_frame(28, data, target, control, patch, ctrl_patch, strip, strip_cell).save(final_png) + render_frame(28, data, target, control, patch, ctrl_patch, strip, strip_cell).save( + final_png + ) print(f"wrote {gif} ({gif.stat().st_size / 1024:.0f} KB, {len(frames)} frames)") print(f"wrote {final_png}") diff --git a/packages/snapcompact/research/snapcompact_r2_filmstrip.py b/packages/snapcompact/research/snapcompact_r2_filmstrip.py index 5f88b31fc..41ffbd022 100755 --- a/packages/snapcompact/research/snapcompact_r2_filmstrip.py +++ b/packages/snapcompact/research/snapcompact_r2_filmstrip.py @@ -40,7 +40,14 @@ EDGE = "#1d2630" DIVERGING = LinearSegmentedColormap.from_list( "carrier_div", - [(0.0, CYAN), (0.30, "#16384a"), (0.50, "#0b1016"), (0.72, "#5c2c18"), (0.90, ORANGE), (1.0, AMBER)], + [ + (0.0, CYAN), + (0.30, "#16384a"), + (0.50, "#0b1016"), + (0.72, "#5c2c18"), + (0.90, ORANGE), + (1.0, AMBER), + ], ) LAYERS = [1, 5, 9, 13, 17, 19, 28] @@ -87,9 +94,14 @@ def sprockets(y: float, x_start: float, x_end: float) -> None: while x + 20 < x_end: ax.add_patch( FancyBboxPatch( - (x, y), 20, 13, + (x, y), + 20, + 13, boxstyle="round,pad=0,rounding_size=4", - facecolor=BG, edgecolor="#27313d", linewidth=1.0, zorder=6, + facecolor=BG, + edgecolor="#27313d", + linewidth=1.0, + zorder=6, ) ) x += 49 @@ -97,8 +109,15 @@ def sprockets(y: float, x_start: float, x_end: float) -> None: def film_band(y0: float, x_end: float) -> None: ax.add_patch( - Rectangle((X0 - 26, y0), x_end - X0 + 26, BAND_H, - facecolor=FILM, edgecolor=EDGE, linewidth=1.2, zorder=2) + Rectangle( + (X0 - 26, y0), + x_end - X0 + 26, + BAND_H, + facecolor=FILM, + edgecolor=EDGE, + linewidth=1.2, + zorder=2, + ) ) sprockets(y0 + 11, X0 - 26, x_end) sprockets(y0 + BAND_H - 24, X0 - 26, x_end) @@ -113,10 +132,29 @@ for y0, label, color in ( (TEXT_BAND_Y, "TEXT REEL", CYAN), (IMAGE_BAND_Y, "IMAGE REEL", ORANGE), ): - ax.text(X0 - 56, y0 + BAND_H / 2, label, color=color, fontsize=13, - fontweight="bold", rotation=90, ha="center", va="center", zorder=8) - ax.text(X0 - 84, y0 + BAND_H / 2, "12 \u00d7 12 carrier cosine", color=MUTED, - fontsize=8, rotation=90, ha="center", va="center", zorder=8) + ax.text( + X0 - 56, + y0 + BAND_H / 2, + label, + color=color, + fontsize=13, + fontweight="bold", + rotation=90, + ha="center", + va="center", + zorder=8, + ) + ax.text( + X0 - 84, + y0 + BAND_H / 2, + "12 \u00d7 12 carrier cosine", + color=MUTED, + fontsize=8, + rotation=90, + ha="center", + va="center", + zorder=8, + ) # ---------------------------------------------------------------- frames VLIM = 0.75 # diagonal (cos=1) clips to amber, off-diagonal structure fills the range @@ -125,12 +163,20 @@ VLIM = 0.75 # diagonal (cos=1) clips to amber, off-diagonal structure fills the def draw_matrix(mat: np.ndarray, cx: float, band_y: float) -> None: x0m, y0m = cx - FS / 2, band_y + 36 ax.imshow( - mat, cmap=DIVERGING, vmin=-VLIM, vmax=VLIM, - extent=(x0m, x0m + FS, y0m + FS, y0m), origin="upper", - interpolation="nearest", zorder=4, + mat, + cmap=DIVERGING, + vmin=-VLIM, + vmax=VLIM, + extent=(x0m, x0m + FS, y0m + FS, y0m), + origin="upper", + interpolation="nearest", + zorder=4, + ) + ax.add_patch( + Rectangle( + (x0m, y0m), FS, FS, fill=False, edgecolor=EDGE, linewidth=1.1, zorder=5 + ) ) - ax.add_patch(Rectangle((x0m, y0m), FS, FS, fill=False, - edgecolor=EDGE, linewidth=1.1, zorder=5)) for i, layer in enumerate(LAYERS): @@ -139,23 +185,62 @@ for i, layer in enumerate(LAYERS): draw_matrix(image_sim[layer], cx, IMAGE_BAND_Y) # frame numbering, film style - ax.text(cx, TEXT_BAND_Y - 12, f"FRAME {i + 1:02d}", color=MUTED, - fontsize=8.5, ha="center", va="bottom", zorder=8) + ax.text( + cx, + TEXT_BAND_Y - 12, + f"FRAME {i + 1:02d}", + color=MUTED, + fontsize=8.5, + ha="center", + va="bottom", + zorder=8, + ) for band_y in (TEXT_BAND_Y, IMAGE_BAND_Y): - ax.text(cx, band_y + 36 + FS + 14, f"LAYER {layer}", color=INK, - fontsize=10, fontweight="bold", ha="center", va="center", zorder=8) + ax.text( + cx, + band_y + 36 + FS + 14, + f"LAYER {layer}", + color=INK, + fontsize=10, + fontweight="bold", + ha="center", + va="center", + zorder=8, + ) # dotted connector between the paired frames - ax.plot([cx, cx], [TEXT_BAND_Y + BAND_H + 3, IMAGE_BAND_Y - 3], - color="#3a4754", linewidth=1.2, linestyle=(0, (1, 3)), zorder=3) + ax.plot( + [cx, cx], + [TEXT_BAND_Y + BAND_H + 3, IMAGE_BAND_Y - 3], + color="#3a4754", + linewidth=1.2, + linestyle=(0, (1, 3)), + zorder=3, + ) # ---------------------------------------------------------------- match meters -ax.text(X0 - 26, METER_Y - 14, "GEOMETRY MATCH", color=INK, fontsize=10, - fontweight="bold", ha="left", va="bottom", zorder=8) -ax.text(X0 + 152, METER_Y - 14, - "amber bar \u2014 RSA: Pearson r of the two reels' off-diagonal structure" - " cyan tick \u2014 matched cross-carrier cosine", - color=MUTED, fontsize=8.5, ha="left", va="bottom", zorder=8) +ax.text( + X0 - 26, + METER_Y - 14, + "GEOMETRY MATCH", + color=INK, + fontsize=10, + fontweight="bold", + ha="left", + va="bottom", + zorder=8, +) +ax.text( + X0 + 152, + METER_Y - 14, + "amber bar \u2014 RSA: Pearson r of the two reels' off-diagonal structure" + " cyan tick \u2014 matched cross-carrier cosine", + color=MUTED, + fontsize=8.5, + ha="left", + va="bottom", + zorder=8, +) BAR_W = FS for i, layer in enumerate(LAYERS): @@ -165,17 +250,51 @@ for i, layer in enumerate(LAYERS): cx = col_cx(i) bx = cx - BAR_W / 2 - ax.add_patch(Rectangle((bx, METER_Y), BAR_W, 12, facecolor=PANEL, - edgecolor=EDGE, linewidth=0.8, zorder=4)) - ax.add_patch(Rectangle((bx, METER_Y), BAR_W * rsa, 12, facecolor=AMBER, - edgecolor="none", zorder=5)) - ax.plot([bx + BAR_W * matched] * 2, [METER_Y - 4, METER_Y + 16], - color=CYAN, linewidth=2.0, zorder=6) + ax.add_patch( + Rectangle( + (bx, METER_Y), + BAR_W, + 12, + facecolor=PANEL, + edgecolor=EDGE, + linewidth=0.8, + zorder=4, + ) + ) + ax.add_patch( + Rectangle( + (bx, METER_Y), BAR_W * rsa, 12, facecolor=AMBER, edgecolor="none", zorder=5 + ) + ) + ax.plot( + [bx + BAR_W * matched] * 2, + [METER_Y - 4, METER_Y + 16], + color=CYAN, + linewidth=2.0, + zorder=6, + ) - ax.text(cx, METER_Y + 32, f"RSA {rsa:.2f}", color=AMBER, fontsize=10.5, - fontweight="bold", ha="center", va="center", zorder=8) - ax.text(cx, METER_Y + 50, f"matched cos {matched:.2f}", color=CYAN, - fontsize=8.5, ha="center", va="center", zorder=8) + ax.text( + cx, + METER_Y + 32, + f"RSA {rsa:.2f}", + color=AMBER, + fontsize=10.5, + fontweight="bold", + ha="center", + va="center", + zorder=8, + ) + ax.text( + cx, + METER_Y + 50, + f"matched cos {matched:.2f}", + color=CYAN, + fontsize=8.5, + ha="center", + va="center", + zorder=8, + ) # ---------------------------------------------------------------- closing callout frame cb_x = X0 + (N_COLS - 1) * CW + 2 @@ -183,21 +302,57 @@ cb_w = X1 - cb_x cb_y0, cb_y1 = TEXT_BAND_Y, METER_Y + METER_H ax.add_patch( FancyBboxPatch( - (cb_x, cb_y0), cb_w, cb_y1 - cb_y0, + (cb_x, cb_y0), + cb_w, + cb_y1 - cb_y0, boxstyle="round,pad=0,rounding_size=10", - facecolor=PANEL, edgecolor=AMBER, linewidth=1.6, zorder=4, + facecolor=PANEL, + edgecolor=AMBER, + linewidth=1.6, + zorder=4, ) ) ccx = cb_x + cb_w / 2 -ax.text(ccx, cb_y0 + 46, "THE SPLICE", color=MUTED, fontsize=10, - ha="center", va="center", zorder=8) -ax.text(ccx, cb_y0 + 122, f"RSA {best['rsa_pearson']:.2f}", color=AMBER, - fontsize=33, fontweight="bold", ha="center", va="center", zorder=8) -ax.text(ccx, cb_y0 + 168, f"@ LAYER {best['layer']}", color=INK, fontsize=14, - fontweight="bold", ha="center", va="center", zorder=8) +ax.text( + ccx, + cb_y0 + 46, + "THE SPLICE", + color=MUTED, + fontsize=10, + ha="center", + va="center", + zorder=8, +) +ax.text( + ccx, + cb_y0 + 122, + f"RSA {best['rsa_pearson']:.2f}", + color=AMBER, + fontsize=33, + fontweight="bold", + ha="center", + va="center", + zorder=8, +) +ax.text( + ccx, + cb_y0 + 168, + f"@ LAYER {best['layer']}", + color=INK, + fontsize=14, + fontweight="bold", + ha="center", + va="center", + zorder=8, +) -ax.plot([cb_x + 28, cb_x + cb_w - 28], [cb_y0 + 206] * 2, - color=EDGE, linewidth=1.0, zorder=5) +ax.plot( + [cb_x + 28, cb_x + cb_w - 28], + [cb_y0 + 206] * 2, + color=EDGE, + linewidth=1.0, + zorder=5, +) facts = [ (f"matched cosine {best['matched_cosine']:.2f}", CYAN), @@ -205,11 +360,29 @@ facts = [ (f"retrieval {int(round(best['match_rank_accuracy'] * 12))}/12", GREEN), ] for j, (line, color) in enumerate(facts): - ax.text(ccx, cb_y0 + 244 + j * 34, line, color=color, fontsize=11.5, - fontweight="bold", ha="center", va="center", zorder=8) + ax.text( + ccx, + cb_y0 + 244 + j * 34, + line, + color=color, + fontsize=11.5, + fontweight="bold", + ha="center", + va="center", + zorder=8, + ) -ax.text(ccx, cb_y0 + 380, "Read it as text or look at\nthe picture \u2014 by layer 19\nthe model files both under\nthe same geometry.", - color=INK, fontsize=10.5, ha="center", va="center", linespacing=1.6, zorder=8) +ax.text( + ccx, + cb_y0 + 380, + "Read it as text or look at\nthe picture \u2014 by layer 19\nthe model files both under\nthe same geometry.", + color=INK, + fontsize=10.5, + ha="center", + va="center", + linespacing=1.6, + zorder=8, +) # the actual L19 splice: the twin pair, miniaturized MINI = 78 @@ -217,36 +390,98 @@ for mat, mx, tag, tcol in ( (text_sim[best["layer"]], ccx - MINI - 9, "text", CYAN), (image_sim[best["layer"]], ccx + 9, "image", ORANGE), ): - ax.imshow(mat, cmap=DIVERGING, vmin=-VLIM, vmax=VLIM, - extent=(mx, mx + MINI, cb_y0 + 444 + MINI, cb_y0 + 444), - origin="upper", interpolation="nearest", zorder=6) - ax.add_patch(Rectangle((mx, cb_y0 + 444), MINI, MINI, fill=False, - edgecolor=EDGE, linewidth=1.0, zorder=7)) - ax.text(mx + MINI / 2, cb_y0 + 444 + MINI + 14, tag, color=tcol, - fontsize=9, ha="center", va="center", zorder=8) -ax.text(ccx, cb_y1 - 36, "two carriers,\none geometry", color=AMBER, fontsize=11, - fontweight="bold", fontstyle="italic", ha="center", va="center", zorder=8) + ax.imshow( + mat, + cmap=DIVERGING, + vmin=-VLIM, + vmax=VLIM, + extent=(mx, mx + MINI, cb_y0 + 444 + MINI, cb_y0 + 444), + origin="upper", + interpolation="nearest", + zorder=6, + ) + ax.add_patch( + Rectangle( + (mx, cb_y0 + 444), + MINI, + MINI, + fill=False, + edgecolor=EDGE, + linewidth=1.0, + zorder=7, + ) + ) + ax.text( + mx + MINI / 2, + cb_y0 + 444 + MINI + 14, + tag, + color=tcol, + fontsize=9, + ha="center", + va="center", + zorder=8, + ) +ax.text( + ccx, + cb_y1 - 36, + "two carriers,\none geometry", + color=AMBER, + fontsize=11, + fontweight="bold", + fontstyle="italic", + ha="center", + va="center", + zorder=8, +) # ---------------------------------------------------------------- title & footer -ax.text(X0 - 26, 64, "TWIN REELS", color=INK, fontsize=34, fontweight="bold", - ha="left", va="center", zorder=8) -ax.text(X0 + 318, 64, "\u2014 the same 12 facts, shot twice", color=AMBER, - fontsize=16, ha="left", va="center", zorder=8) ax.text( - X0 - 26, 118, + X0 - 26, + 64, + "TWIN REELS", + color=INK, + fontsize=34, + fontweight="bold", + ha="left", + va="center", + zorder=8, +) +ax.text( + X0 + 318, + 64, + "\u2014 the same 12 facts, shot twice", + color=AMBER, + fontsize=16, + ha="left", + va="center", + zorder=8, +) +ax.text( + X0 - 26, + 118, "Twelve question\u2013answer pairs enter Qwen2.5-VL-7B twice: once as text, once rendered into pixels. " "Each frame is the 12\u00d712 cosine similarity between carrier states at one layer \u2014 " "the two reels print the same relational structure from the very first frames.", - color=MUTED, fontsize=11.5, ha="left", va="center", zorder=8, + color=MUTED, + fontsize=11.5, + ha="left", + va="center", + zorder=8, ) ax.text( - X0 - 26, FOOT_Y, + X0 - 26, + FOOT_Y, "data: results/qwen-carrier-convergence-n12 (carrier_convergence.npz \u00b7 summary.json) \u00b7 " "carrier-centered cosine of hidden states, d = 3584, 29 layers \u00b7 " "RSA = Pearson r over the 66 off-diagonal pairs \u00b7 " - "diverging scale \u2212%.2f \u2026 +%.2f (cyan \u2192 dark \u2192 orange)" % (VLIM, VLIM), - color=MUTED, fontsize=9, ha="left", va="center", zorder=8, + "diverging scale \u2212%.2f \u2026 +%.2f (cyan \u2192 dark \u2192 orange)" + % (VLIM, VLIM), + color=MUTED, + fontsize=9, + ha="left", + va="center", + zorder=8, ) # ---------------------------------------------------------------- save diff --git a/packages/snapcompact/research/snapcompact_r2_hero.py b/packages/snapcompact/research/snapcompact_r2_hero.py index 49370267d..1f04be977 100755 --- a/packages/snapcompact/research/snapcompact_r2_hero.py +++ b/packages/snapcompact/research/snapcompact_r2_hero.py @@ -98,14 +98,24 @@ def body_font(size: float) -> ImageFont.FreeTypeFont: def mono_font(size: float) -> ImageFont.FreeTypeFont: - for path in ["/System/Library/Fonts/Monaco.ttf", "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf"]: + for path in [ + "/System/Library/Fonts/Monaco.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf", + ]: f = font_at(path, size) if f is not None: return f return ImageFont.load_default() -def tracked(draw: ImageDraw.ImageDraw, xy: tuple[int, int], text: str, font, fill, tracking: float = 0.0) -> int: +def tracked( + draw: ImageDraw.ImageDraw, + xy: tuple[int, int], + text: str, + font, + fill, + tracking: float = 0.0, +) -> int: """Draw text with letterspacing; returns end x.""" x, y = xy t = u(tracking) @@ -115,7 +125,9 @@ def tracked(draw: ImageDraw.ImageDraw, xy: tuple[int, int], text: str, font, fil return int(x) -def tracked_width(draw: ImageDraw.ImageDraw, text: str, font, tracking: float = 0.0) -> float: +def tracked_width( + draw: ImageDraw.ImageDraw, text: str, font, tracking: float = 0.0 +) -> float: t = u(tracking) return sum(draw.textlength(ch, font=font) + t for ch in text) - (t if text else 0) @@ -131,9 +143,15 @@ def bezier(p0, p1, p2, n=64): def load_data(): - conv = json.loads((ROOT / "results" / "qwen-carrier-convergence-n12" / "summary.json").read_text()) - entry = json.loads((ROOT / "results" / "qwen-token-entry-q3" / "token_entry.json").read_text()) - carrier = Image.open(ROOT / "results" / "qwen-logit-lens-q3" / "images" / "image-carrier.png").convert("RGB") + conv = json.loads( + (ROOT / "results" / "qwen-carrier-convergence-n12" / "summary.json").read_text() + ) + entry = json.loads( + (ROOT / "results" / "qwen-token-entry-q3" / "token_entry.json").read_text() + ) + carrier = Image.open( + ROOT / "results" / "qwen-logit-lens-q3" / "images" / "image-carrier.png" + ).convert("RGB") if carrier.size != (1568, 1568): carrier = carrier.resize((1568, 1568), Image.LANCZOS) @@ -149,14 +167,22 @@ def load_data(): "n": n_q, } assert stats["layer"] == 19 and abs(stats["matched"] - 0.66) < 0.01 - assert abs(stats["rsa"] - 0.85) < 0.01 and stats["retrieved"] == 12 and stats["n"] == 12 + assert ( + abs(stats["rsa"] - 0.85) < 0.01 + and stats["retrieved"] == 12 + and stats["n"] == 12 + ) toks = {t["i"]: t for t in entry["tokens"]} answer = [t for t in entry["tokens"] if t["answer"]] assert [t["str"] for t in answer] == ["spect", "acular"] assert [t["id"] for t in answer] == [67082, 23006] - ctx_before = "…" + "".join(toks[i]["str"] for i in range(23, 32)) # " make the 50th Super Bowl \"" - ctx_after = "".join(toks[i]["str"] for i in range(34, 39)) + "…" # "\" and that it would" + ctx_before = "…" + "".join( + toks[i]["str"] for i in range(23, 32) + ) # " make the 50th Super Bowl \"" + ctx_after = ( + "".join(toks[i]["str"] for i in range(34, 39)) + "…" + ) # "\" and that it would" grid = entry["image_grid"] # 56 word_idx = entry["image_answer_token_indices"][:4] # [310, 311, 312, 313] @@ -225,7 +251,12 @@ def additive_base() -> np.ndarray: n = 9 for k in range(n): f = k / (n - 1) - y0 = PANEL_TOP + 120 + f * (PANEL_BOT - PANEL_TOP - 240) + rng.uniform(-14, 14) + y0 = ( + PANEL_TOP + + 120 + + f * (PANEL_BOT - PANEL_TOP - 240) + + rng.uniform(-14, 14) + ) ang = (f - 0.5) * 1.45 + rng.uniform(-0.07, 0.07) r = 168 x2 = CORE[0] - x_sign * r * math.cos(ang) @@ -242,7 +273,10 @@ def additive_base() -> np.ndarray: i = int(t * len(pts)) px, py = u(pts[i][0]), u(pts[i][1]) rr = u(3.2) - gd.ellipse([px - rr, py - rr, px + rr, py + rr], fill=tuple(int(c * fade) for c in color)) + gd.ellipse( + [px - rr, py - rr, px + rr, py + rr], + fill=tuple(int(c * fade) for c in color), + ) streams(LP[2], 1.0, AMBER) streams(RP[0], -1.0, CYAN) @@ -255,12 +289,16 @@ def additive_base() -> np.ndarray: layer = Image.new("RGB", (W, H), (0, 0, 0)) ld = ImageDraw.Draw(layer) tw = ld.textlength(CORE_WORD, font=f) - ld.text((u(CORE[0]) - tw / 2, u(CORE[1] - 54)), CORE_WORD, font=f, fill=(255, 232, 170)) + ld.text( + (u(CORE[0]) - tw / 2, u(CORE[1] - 54)), CORE_WORD, font=f, fill=(255, 232, 170) + ) img += np.asarray(layer.filter(ImageFilter.GaussianBlur(u(9))), np.float32) * 0.8 return img -def head_bars(ov: ImageDraw.ImageDraw, x: float, y_mid: float, values, color, label: str): +def head_bars( + ov: ImageDraw.ImageDraw, x: float, y_mid: float, values, color, label: str +): """Tiny bar strip of a real 10-dim vector head, centered on its axis.""" vmax = max(abs(v) for v in values) bw, gap, amp = 24, 11, 26 @@ -281,11 +319,24 @@ def draw_left_panel(ov: ImageDraw.ImageDraw, answer, ctx, counts, heads): pad = 44 ctx_before, ctx_after = ctx - tracked(ov, (u(x0 + pad), u(y0 + 34)), "TEXT CARRIER", label_font(30), AMBER, tracking=5) + tracked( + ov, (u(x0 + pad), u(y0 + 34)), "TEXT CARRIER", label_font(30), AMBER, tracking=5 + ) sub = f"{counts['text_tokens']:,} BPE TOKENS" f_sub = label_font(21) - tracked(ov, (int(u(x1 - pad) - tracked_width(ov, sub, f_sub, 2)), u(y0 + 42)), sub, f_sub, MUTED, tracking=2) - ov.line([u(x0 + pad), u(y0 + 88), u(x1 - pad), u(y0 + 88)], fill=(*DIVIDER, 255), width=u(1.2)) + tracked( + ov, + (int(u(x1 - pad) - tracked_width(ov, sub, f_sub, 2)), u(y0 + 42)), + sub, + f_sub, + MUTED, + tracking=2, + ) + ov.line( + [u(x0 + pad), u(y0 + 88), u(x1 - pad), u(y0 + 88)], + fill=(*DIVIDER, 255), + width=u(1.2), + ) f_ctx = mono_font(23) ov.text((u(x0 + pad), u(y0 + 116)), ctx_before, font=f_ctx, fill=(110, 118, 126)) @@ -306,7 +357,9 @@ def draw_left_panel(ov: ImageDraw.ImageDraw, answer, ctx, counts, heads): width=u(1.6), ) ov.text((u(px + 20), u(py + 14)), s, font=f_tok, fill=(255, 224, 150)) - ov.text((u(px + 20), u(py + 110)), f"id {t['id']}", font=f_id, fill=(196, 156, 72)) + ov.text( + (u(px + 20), u(py + 110)), f"id {t['id']}", font=f_id, fill=(196, 156, 72) + ) px += wpx + 40 + 22 ov.text((u(x0 + pad), u(y0 + 330)), ctx_after, font=f_ctx, fill=(110, 118, 126)) @@ -323,7 +376,12 @@ def draw_left_panel(ov: ImageDraw.ImageDraw, answer, ctx, counts, heads): fy = y0 + 500 f_fact = body_font(23) - ov.text((u(x0 + pad), u(fy)), f"{counts['chars']:,} characters of one SQuAD passage,", font=f_fact, fill=MUTED) + ov.text( + (u(x0 + pad), u(fy)), + f"{counts['chars']:,} characters of one SQuAD passage,", + font=f_fact, + fill=MUTED, + ) ov.text( (u(x0 + pad), u(fy + 36)), f"tokenized into {counts['text_tokens']:,} ids, each a {counts['embed_dim']:,}-dim row", @@ -332,15 +390,35 @@ def draw_left_panel(ov: ImageDraw.ImageDraw, answer, ctx, counts, heads): ) -def draw_right_panel(base_img: Image.Image, ov: ImageDraw.ImageDraw, carrier: Image.Image, word_idx, counts, heads): +def draw_right_panel( + base_img: Image.Image, + ov: ImageDraw.ImageDraw, + carrier: Image.Image, + word_idx, + counts, + heads, +): x0, y0, x1, _ = RP pad = 44 - tracked(ov, (u(x0 + pad), u(y0 + 34)), "IMAGE CARRIER", label_font(30), CYAN, tracking=5) + tracked( + ov, (u(x0 + pad), u(y0 + 34)), "IMAGE CARRIER", label_font(30), CYAN, tracking=5 + ) sub = f"{counts['image_tokens']:,} PATCHES" f_sub = label_font(21) - tracked(ov, (int(u(x1 - pad) - tracked_width(ov, sub, f_sub, 2)), u(y0 + 42)), sub, f_sub, MUTED, tracking=2) - ov.line([u(x0 + pad), u(y0 + 88), u(x1 - pad), u(y0 + 88)], fill=(*DIVIDER, 255), width=u(1.2)) + tracked( + ov, + (int(u(x1 - pad) - tracked_width(ov, sub, f_sub, 2)), u(y0 + 42)), + sub, + f_sub, + MUTED, + tracking=2, + ) + ov.line( + [u(x0 + pad), u(y0 + 88), u(x1 - pad), u(y0 + 88)], + fill=(*DIVIDER, 255), + width=u(1.2), + ) # Crop: patch rows 4..8, cols 27..38 of the 56x56 grid (28px cells). pp = counts["patch_px"] @@ -357,18 +435,32 @@ def draw_right_panel(base_img: Image.Image, ov: ImageDraw.ImageDraw, carrier: Im cell = pp * scale # 56 in 2400-space grid_color = (CYAN[0], CYAN[1], CYAN[2], 46) for c in range(c1 - c0 + 1): - ov.line([u(bx + c * cell), u(by), u(bx + c * cell), u(by + disp_h)], fill=grid_color, width=u(1)) + ov.line( + [u(bx + c * cell), u(by), u(bx + c * cell), u(by + disp_h)], + fill=grid_color, + width=u(1), + ) for r in range(r1 - r0 + 1): - ov.line([u(bx), u(by + r * cell), u(bx + disp_w), u(by + r * cell)], fill=grid_color, width=u(1)) + ov.line( + [u(bx), u(by + r * cell), u(bx + disp_w), u(by + r * cell)], + fill=grid_color, + width=u(1), + ) # Highlight the four answer patches (grid row 5, cols 30..33) as one run. grid = counts["grid"] gr, gc = word_idx[0] // grid, word_idx[0] % grid hx, hy = bx + (gc - c0) * cell, by + (gr - r0) * cell hw = len(word_idx) * cell - ov.rectangle([u(hx), u(hy), u(hx + hw), u(hy + cell)], outline=(*CYAN, 240), width=u(2.4)) + ov.rectangle( + [u(hx), u(hy), u(hx + hw), u(hy + cell)], outline=(*CYAN, 240), width=u(2.4) + ) for k in range(1, len(word_idx)): - ov.line([u(hx + k * cell), u(hy), u(hx + k * cell), u(hy + cell)], fill=(*CYAN, 130), width=u(1.2)) + ov.line( + [u(hx + k * cell), u(hy), u(hx + k * cell), u(hy + cell)], + fill=(*CYAN, 130), + width=u(1.2), + ) ov.text( (u(bx), u(by + disp_h + 18)), @@ -389,7 +481,12 @@ def draw_right_panel(base_img: Image.Image, ov: ImageDraw.ImageDraw, carrier: Im fy = y0 + 500 f_fact = body_font(23) - ov.text((u(bx), u(fy)), "the same passage, rendered to a 1568 × 1568 px bitmap,", font=f_fact, fill=MUTED) + ov.text( + (u(bx), u(fy)), + "the same passage, rendered to a 1568 × 1568 px bitmap,", + font=f_fact, + fill=MUTED, + ) ov.text( (u(bx), u(fy + 36)), f"seen as {counts['image_tokens']:,} patches of {counts['patch_px']} px, each a {counts['visual_dim']:,}-dim vector", @@ -427,7 +524,14 @@ def draw_core(ov: ImageDraw.ImageDraw, stats): cap = f"BY LAYER {stats['layer']} OF {stats['n_layers'] - 1}, ONE SHARED STATE" f_c = label_font(23) cw = tracked_width(ov, cap, f_c, 4) - tracked(ov, (int(u(CORE[0]) - cw / 2), u(CORE[1] + 96)), cap, f_c, (228, 222, 196), tracking=4) + tracked( + ov, + (int(u(CORE[0]) - cw / 2), u(CORE[1] + 96)), + cap, + f_c, + (228, 222, 196), + tracking=4, + ) def draw_stats_strip(ov: ImageDraw.ImageDraw, stats): @@ -460,7 +564,9 @@ def rounded_panel(overlay: ImageDraw.ImageDraw, box, accent, alpha_fill=216): x0, y0, x1, y1 = (u(v) for v in box) r = u(22) overlay.rounded_rectangle([x0, y0, x1, y1], radius=r, fill=(*PANEL, alpha_fill)) - overlay.rounded_rectangle([x0, y0, x1, y1], radius=r, outline=(*accent, 70), width=u(1.4)) + overlay.rounded_rectangle( + [x0, y0, x1, y1], radius=r, outline=(*accent, 70), width=u(1.4) + ) def main() -> None: diff --git a/packages/snapcompact/research/snapcompact_r2_metro.py b/packages/snapcompact/research/snapcompact_r2_metro.py index d7da1477a..190dac6c6 100755 --- a/packages/snapcompact/research/snapcompact_r2_metro.py +++ b/packages/snapcompact/research/snapcompact_r2_metro.py @@ -25,7 +25,9 @@ import matplotlib.pyplot as plt # noqa: E402 from matplotlib.patches import Circle # noqa: E402 HERE = os.path.dirname(os.path.abspath(__file__)) -SUMMARY_PATH = os.path.join(HERE, "results", "qwen-carrier-convergence-n12", "summary.json") +SUMMARY_PATH = os.path.join( + HERE, "results", "qwen-carrier-convergence-n12", "summary.json" +) LENS_PATH = os.path.join(HERE, "results", "qwen-logit-lens-q3", "logit_lens.json") OUT_DIR = os.path.join(HERE, "results", "agent-r2-metro") OUT_PNG = os.path.join(OUT_DIR, "metro.png") @@ -150,20 +152,51 @@ def main(): ax.plot([i, i], [-3.55, 3.95], color=GRID, lw=1.0, zorder=1) # convergence axis - ax.plot([-1.5, 28.0], [0, 0], color=MUTED, lw=1.0, alpha=0.28, - linestyle=(0, (1, 3)), zorder=1) - ax.text(8.0, 0.16, "convergence axis", color=MUTED, alpha=0.55, - fontsize=8.5, style="italic", family=SANS, ha="left", zorder=2) + ax.plot( + [-1.5, 28.0], + [0, 0], + color=MUTED, + lw=1.0, + alpha=0.28, + linestyle=(0, (1, 3)), + zorder=1, + ) + ax.text( + 8.0, + 0.16, + "convergence axis", + color=MUTED, + alpha=0.55, + fontsize=8.5, + style="italic", + family=SANS, + ha="left", + zorder=2, + ) # ---- metro lines: glow, casing, stroke ------------------------------- for path, col in ((path_t, CYAN), (path_i, ORANGE)): px, py = path[:, 0], path[:, 1] ax.plot(px, py, color=col, lw=24, alpha=0.05, solid_capstyle="round", zorder=2) ax.plot(px, py, color=col, lw=16, alpha=0.07, solid_capstyle="round", zorder=2) - ax.plot(px, py, color=BG, lw=14, solid_capstyle="round", - solid_joinstyle="round", zorder=3) - ax.plot(px, py, color=col, lw=9.5, solid_capstyle="round", - solid_joinstyle="round", zorder=4) + ax.plot( + px, + py, + color=BG, + lw=14, + solid_capstyle="round", + solid_joinstyle="round", + zorder=3, + ) + ax.plot( + px, + py, + color=col, + lw=9.5, + solid_capstyle="round", + solid_joinstyle="round", + zorder=4, + ) # ---- stations --------------------------------------------------------- named = {19, 24, 27, 28} @@ -171,54 +204,141 @@ def main(): for y, col in ((y_text[i], CYAN), (y_img[i], ORANGE)): if i in named: continue - ax.scatter([i], [y], s=115, facecolor=BG, edgecolor=col, - linewidths=2.1, zorder=6) + ax.scatter( + [i], [y], s=115, facecolor=BG, edgecolor=col, linewidths=2.1, zorder=6 + ) # L19 interchange capsule (the two lines meet in one station) cap_lw_outer = 0.56 * pt_per_unit - ax.plot([19, 19], [y_img[19], y_text[19]], color=INK, - lw=cap_lw_outer, solid_capstyle="round", zorder=5) - ax.plot([19, 19], [y_img[19], y_text[19]], color=PANEL, - lw=cap_lw_outer - 7.5, solid_capstyle="round", zorder=5) - ax.scatter([19, 19], [y_text[19], y_img[19]], s=92, - c=[CYAN, ORANGE], edgecolor=BG, linewidths=1.2, zorder=6) + ax.plot( + [19, 19], + [y_img[19], y_text[19]], + color=INK, + lw=cap_lw_outer, + solid_capstyle="round", + zorder=5, + ) + ax.plot( + [19, 19], + [y_img[19], y_text[19]], + color=PANEL, + lw=cap_lw_outer - 7.5, + solid_capstyle="round", + zorder=5, + ) + ax.scatter( + [19, 19], + [y_text[19], y_img[19]], + s=92, + c=[CYAN, ORANGE], + edgecolor=BG, + linewidths=1.2, + zorder=6, + ) # L24 interchange ring on the image line (pixels decode to vocabulary) - ax.scatter([24], [y_img[24]], s=300, facecolor=PANEL, edgecolor=INK, - linewidths=2.8, zorder=6) + ax.scatter( + [24], + [y_img[24]], + s=300, + facecolor=PANEL, + edgecolor=INK, + linewidths=2.8, + zorder=6, + ) ax.scatter([24], [y_img[24]], s=58, facecolor=ORANGE, edgecolor="none", zorder=6) - ax.scatter([24], [y_text[24]], s=115, facecolor=BG, edgecolor=CYAN, - linewidths=2.1, zorder=6) + ax.scatter( + [24], + [y_text[24]], + s=115, + facecolor=BG, + edgecolor=CYAN, + linewidths=2.1, + zorder=6, + ) # L27 white-ring stations on both lines (terminal approach) for y, col in ((y_text[27], CYAN), (y_img[27], ORANGE)): - ax.scatter([27], [y], s=170, facecolor=PANEL, edgecolor=INK, - linewidths=2.3, zorder=6) + ax.scatter( + [27], [y], s=170, facecolor=PANEL, edgecolor=INK, linewidths=2.3, zorder=6 + ) ax.scatter([27], [y], s=34, facecolor=col, edgecolor="none", zorder=6) # L28 terminus: double ring over both tracks - ax.add_patch(Circle((28, 0), 1.02, facecolor=PANEL, edgecolor=INK, - lw=3.2, zorder=5)) - ax.add_patch(Circle((28, 0), 0.66, facecolor="none", edgecolor=INK, - lw=1.3, alpha=0.85, zorder=5)) - ax.scatter([27.78, 28.22], [0, 0], s=120, c=[CYAN, ORANGE], - edgecolor=BG, linewidths=1.4, zorder=6) - ax.text(28, -0.42, "TERMINUS", color=MUTED, fontsize=6.8, family=MONO, - ha="center", va="center", zorder=7) + ax.add_patch( + Circle((28, 0), 1.02, facecolor=PANEL, edgecolor=INK, lw=3.2, zorder=5) + ) + ax.add_patch( + Circle( + (28, 0), 0.66, facecolor="none", edgecolor=INK, lw=1.3, alpha=0.85, zorder=5 + ) + ) + ax.scatter( + [27.78, 28.22], + [0, 0], + s=120, + c=[CYAN, ORANGE], + edgecolor=BG, + linewidths=1.4, + zorder=6, + ) + ax.text( + 28, + -0.42, + "TERMINUS", + color=MUTED, + fontsize=6.8, + family=MONO, + ha="center", + va="center", + zorder=7, + ) # ---- carrier labels (depots) ------------------------------------------ geo = d["geometry"] - ax.text(-1.55, y_text[0] + 0.95, "TEXT CARRIER", color=CYAN, fontsize=12.5, - family=SANS, fontweight="bold", ha="left", zorder=7) - ax.text(-1.55, y_text[0] + 0.48, - f"the page as typed tokens \u00b7 {geo['capacity']:,} chars", - color=MUTED, fontsize=9, family=SANS, ha="left", zorder=7) - ax.text(-1.55, y_img[0] - 0.62, "IMAGE CARRIER", color=ORANGE, fontsize=12.5, - family=SANS, fontweight="bold", ha="left", zorder=7) - ax.text(-1.55, y_img[0] - 1.09, - f"the same page as a {d['size_px']} px bitmap \u00b7 " - f"{geo['cols']}\u00d7{geo['rows']} cell grid", - color=MUTED, fontsize=9, family=SANS, ha="left", zorder=7) + ax.text( + -1.55, + y_text[0] + 0.95, + "TEXT CARRIER", + color=CYAN, + fontsize=12.5, + family=SANS, + fontweight="bold", + ha="left", + zorder=7, + ) + ax.text( + -1.55, + y_text[0] + 0.48, + f"the page as typed tokens \u00b7 {geo['capacity']:,} chars", + color=MUTED, + fontsize=9, + family=SANS, + ha="left", + zorder=7, + ) + ax.text( + -1.55, + y_img[0] - 0.62, + "IMAGE CARRIER", + color=ORANGE, + fontsize=12.5, + family=SANS, + fontweight="bold", + ha="left", + zorder=7, + ) + ax.text( + -1.55, + y_img[0] - 1.09, + f"the same page as a {d['size_px']} px bitmap \u00b7 " + f"{geo['cols']}\u00d7{geo['rows']} cell grid", + color=MUTED, + fontsize=9, + family=SANS, + ha="left", + zorder=7, + ) # ---- named-station callouts ------------------------------------------- def leader(x, y_from, y_to, color=MUTED, alpha=0.65): @@ -226,117 +346,379 @@ def main(): # L1: instant alignment leader(1, y_img[1] - 0.18, -2.18) - ax.text(1.7, -2.35, "L1 \u00b7 INSTANT ALIGNMENT", color=INK, fontsize=10.5, - family=SANS, fontweight="bold", ha="left", zorder=7) - ax.text(1.7, -2.78, - f"matched cos {fmt2(cos[1])} \u00b7 RSA {fmt2(d['rsa'][1])}", - color=MUTED, fontsize=8.8, family=MONO, ha="left", zorder=7) - ax.text(1.7, -3.14, - f"retrieval {int(round(d['acc'][1] * 12))}/12 \u2014 12/12 from L2 onward", - color=MUTED, fontsize=8.8, family=MONO, ha="left", zorder=7) + ax.text( + 1.7, + -2.35, + "L1 \u00b7 INSTANT ALIGNMENT", + color=INK, + fontsize=10.5, + family=SANS, + fontweight="bold", + ha="left", + zorder=7, + ) + ax.text( + 1.7, + -2.78, + f"matched cos {fmt2(cos[1])} \u00b7 RSA {fmt2(d['rsa'][1])}", + color=MUTED, + fontsize=8.8, + family=MONO, + ha="left", + zorder=7, + ) + ax.text( + 1.7, + -3.14, + f"retrieval {int(round(d['acc'][1] * 12))}/12 \u2014 12/12 from L2 onward", + color=MUTED, + fontsize=8.8, + family=MONO, + ha="left", + zorder=7, + ) # L13: first close pass leader(13, y_text[13] + 0.18, 1.62) - ax.text(13, 1.84, f"L13 \u00b7 first close pass \u00b7 cos {fmt2(cos[13])}", - color=MUTED, fontsize=8.8, family=MONO, ha="center", zorder=7) + ax.text( + 13, + 1.84, + f"L13 \u00b7 first close pass \u00b7 cos {fmt2(cos[13])}", + color=MUTED, + fontsize=8.8, + family=MONO, + ha="center", + zorder=7, + ) # L19: geometry locks (star station) leader(19, y_text[19] + 0.62, 2.42, color=AMBER, alpha=0.8) - ax.text(19, 3.42, "L19 \u00b7 GEOMETRY LOCKS", color=AMBER, fontsize=14, - family=SANS, fontweight="bold", ha="center", zorder=7) - ax.text(19, 2.96, - f"matched cos {fmt2(cos[19])} \u00b7 mismatched {fmt2(d['mism'][19])}", - color=INK, fontsize=9.6, family=MONO, ha="center", zorder=7) - ax.text(19, 2.58, - f"RSA {fmt2(d['rsa'][19])} \u00b7 retrieval 12/12 \u2014 closest approach", - color=MUTED, fontsize=9.6, family=MONO, ha="center", zorder=7) + ax.text( + 19, + 3.42, + "L19 \u00b7 GEOMETRY LOCKS", + color=AMBER, + fontsize=14, + family=SANS, + fontweight="bold", + ha="center", + zorder=7, + ) + ax.text( + 19, + 2.96, + f"matched cos {fmt2(cos[19])} \u00b7 mismatched {fmt2(d['mism'][19])}", + color=INK, + fontsize=9.6, + family=MONO, + ha="center", + zorder=7, + ) + ax.text( + 19, + 2.58, + f"RSA {fmt2(d['rsa'][19])} \u00b7 retrieval 12/12 \u2014 closest approach", + color=MUTED, + fontsize=9.6, + family=MONO, + ha="center", + zorder=7, + ) # L23: small drift leader(23, y_text[23] + 0.18, 1.30) - ax.text(23, 1.52, f"L23 \u00b7 small drift \u00b7 cos {fmt2(cos[23])}", - color=MUTED, fontsize=8.8, family=MONO, ha="center", zorder=7) + ax.text( + 23, + 1.52, + f"L23 \u00b7 small drift \u00b7 cos {fmt2(cos[23])}", + color=MUTED, + fontsize=8.8, + family=MONO, + ha="center", + zorder=7, + ) # L24: pixels decode to vocabulary leader(24, y_img[24] - 0.32, -1.92, color=ORANGE, alpha=0.8) - ax.text(24, -2.18, "L24 \u00b7 PIXELS DECODE TO VOCABULARY", color=ORANGE, - fontsize=12.5, family=SANS, fontweight="bold", ha="center", zorder=7) - ax.text(24, -2.62, - f"visual tok[310] top-1 \u2192 'acular' \u00b7 p {d['p24']:.2f}", - color=INK, fontsize=9.4, family=MONO, ha="center", zorder=7) - ax.text(24, -3.00, - f"rising to p {d['p28']:.2f} by L28 \u2014 " - "the answer's second BPE piece", - color=MUTED, fontsize=9.4, family=MONO, ha="center", zorder=7) + ax.text( + 24, + -2.18, + "L24 \u00b7 PIXELS DECODE TO VOCABULARY", + color=ORANGE, + fontsize=12.5, + family=SANS, + fontweight="bold", + ha="center", + zorder=7, + ) + ax.text( + 24, + -2.62, + f"visual tok[310] top-1 \u2192 'acular' \u00b7 p {d['p24']:.2f}", + color=INK, + fontsize=9.4, + family=MONO, + ha="center", + zorder=7, + ) + ax.text( + 24, + -3.00, + f"rising to p {d['p28']:.2f} by L28 \u2014 the answer's second BPE piece", + color=MUTED, + fontsize=9.4, + family=MONO, + ha="center", + zorder=7, + ) # L27-L28 terminal (block above the terminus circle) leader(28, 1.18, 1.86, color=AMBER, alpha=0.8) - ax.text(28, 3.00, "L27\u2013L28 \u00b7 TERMINAL", color=INK, fontsize=12.5, - family=SANS, fontweight="bold", ha="center", zorder=7) - ax.text(28, 2.56, f"SAME ANSWER: \u201c{d['answer']}\u201d", color=AMBER, - fontsize=10.5, family=SANS, fontweight="bold", ha="center", zorder=7) - ax.text(28, 2.18, - f"matched cos {fmt2(cos[27])} \u2192 {fmt2(cos[28])} \u00b7 retrieval 12/12", - color=MUTED, fontsize=8.8, family=MONO, ha="center", zorder=7) + ax.text( + 28, + 3.00, + "L27\u2013L28 \u00b7 TERMINAL", + color=INK, + fontsize=12.5, + family=SANS, + fontweight="bold", + ha="center", + zorder=7, + ) + ax.text( + 28, + 2.56, + f"SAME ANSWER: \u201c{d['answer']}\u201d", + color=AMBER, + fontsize=10.5, + family=SANS, + fontweight="bold", + ha="center", + zorder=7, + ) + ax.text( + 28, + 2.18, + f"matched cos {fmt2(cos[27])} \u2192 {fmt2(cos[28])} \u00b7 retrieval 12/12", + color=MUTED, + fontsize=8.8, + family=MONO, + ha="center", + zorder=7, + ) # ---- station index + matched-cosine gauge rows ------------------------- hl = {19: AMBER, 24: ORANGE, 27: INK, 28: INK} - ax.text(-0.55, -4.45, "layer", color=MUTED, fontsize=8, style="italic", - family=SANS, ha="right", va="center", zorder=7) - ax.text(-0.55, -5.02, "matched cos", color=MUTED, fontsize=8, style="italic", - family=SANS, ha="right", va="center", zorder=7) + ax.text( + -0.55, + -4.45, + "layer", + color=MUTED, + fontsize=8, + style="italic", + family=SANS, + ha="right", + va="center", + zorder=7, + ) + ax.text( + -0.55, + -5.02, + "matched cos", + color=MUTED, + fontsize=8, + style="italic", + family=SANS, + ha="right", + va="center", + zorder=7, + ) for i in range(n_layers): col = hl.get(i, MUTED) w = "bold" if i in hl else "normal" - ax.text(i, -4.45, f"L{i}", color=col, fontsize=7.6, family=MONO, - ha="center", va="center", fontweight=w, zorder=7) + ax.text( + i, + -4.45, + f"L{i}", + color=col, + fontsize=7.6, + family=MONO, + ha="center", + va="center", + fontweight=w, + zorder=7, + ) val = f"{cos[i]:.2f}".lstrip("0") - ax.text(i, -5.02, val, color=col, fontsize=7.6, family=MONO, - ha="center", va="center", fontweight=w, zorder=7) + ax.text( + i, + -5.02, + val, + color=col, + fontsize=7.6, + family=MONO, + ha="center", + va="center", + fontweight=w, + zorder=7, + ) # ---- title ------------------------------------------------------------- - ax.text(-1.9, 8.05, "THE CONVERGENCE LINE", color=INK, fontsize=29, - family=SANS, fontweight="bold", ha="left", va="top", zorder=7) - ax.text(-1.9, 6.92, - "One Wikipedia page, two carriers: typed tokens (cyan) and a " - f"{d['size_px']} px screenshot (orange) ride Qwen2.5-VL-7B's 29 decoder layers.", - color=MUTED, fontsize=12, family=SANS, ha="left", va="top", zorder=7) - ax.text(-1.9, 6.42, - "The closer the tracks, the more the two internal representations agree " - f"\u2014 track gap \u221d 1 {MINUS} matched cosine, n = {d['n_q']} questions.", - color=MUTED, fontsize=12, family=SANS, ha="left", va="top", zorder=7) + ax.text( + -1.9, + 8.05, + "THE CONVERGENCE LINE", + color=INK, + fontsize=29, + family=SANS, + fontweight="bold", + ha="left", + va="top", + zorder=7, + ) + ax.text( + -1.9, + 6.92, + "One Wikipedia page, two carriers: typed tokens (cyan) and a " + f"{d['size_px']} px screenshot (orange) ride Qwen2.5-VL-7B's 29 decoder layers.", + color=MUTED, + fontsize=12, + family=SANS, + ha="left", + va="top", + zorder=7, + ) + ax.text( + -1.9, + 6.42, + "The closer the tracks, the more the two internal representations agree " + f"\u2014 track gap \u221d 1 {MINUS} matched cosine, n = {d['n_q']} questions.", + color=MUTED, + fontsize=12, + family=SANS, + ha="left", + va="top", + zorder=7, + ) # ---- legend (top right) ------------------------------------------------- lx = 22.9 - ax.plot([lx, lx + 1.3], [7.95, 7.95], color=CYAN, lw=8, - solid_capstyle="round", zorder=7) - ax.text(lx + 1.65, 7.95, "TEXT CARRIER", color=INK, fontsize=10, - family=SANS, fontweight="bold", ha="left", va="center", zorder=7) - ax.plot([lx, lx + 1.3], [7.32, 7.32], color=ORANGE, lw=8, - solid_capstyle="round", zorder=7) - ax.text(lx + 1.65, 7.32, "IMAGE CARRIER", color=INK, fontsize=10, - family=SANS, fontweight="bold", ha="left", va="center", zorder=7) - ax.text(lx, 6.62, f"track gap \u221d 1 {MINUS} matched cosine(text, image)", - color=MUTED, fontsize=9, family=MONO, ha="left", va="center", zorder=7) + ax.plot( + [lx, lx + 1.3], [7.95, 7.95], color=CYAN, lw=8, solid_capstyle="round", zorder=7 + ) + ax.text( + lx + 1.65, + 7.95, + "TEXT CARRIER", + color=INK, + fontsize=10, + family=SANS, + fontweight="bold", + ha="left", + va="center", + zorder=7, + ) + ax.plot( + [lx, lx + 1.3], + [7.32, 7.32], + color=ORANGE, + lw=8, + solid_capstyle="round", + zorder=7, + ) + ax.text( + lx + 1.65, + 7.32, + "IMAGE CARRIER", + color=INK, + fontsize=10, + family=SANS, + fontweight="bold", + ha="left", + va="center", + zorder=7, + ) + ax.text( + lx, + 6.62, + f"track gap \u221d 1 {MINUS} matched cosine(text, image)", + color=MUTED, + fontsize=9, + family=MONO, + ha="left", + va="center", + zorder=7, + ) # wide pair = L0 - ax.plot([lx, lx + 1.0], [6.18, 6.18], color=CYAN, lw=4, solid_capstyle="round", zorder=7) - ax.plot([lx, lx + 1.0], [5.74, 5.74], color=ORANGE, lw=4, solid_capstyle="round", zorder=7) - ax.text(lx + 1.65, 5.96, f"cos {fmt2(cos[0])} \u2014 far apart (L0)", - color=MUTED, fontsize=9, family=MONO, ha="left", va="center", zorder=7) + ax.plot( + [lx, lx + 1.0], [6.18, 6.18], color=CYAN, lw=4, solid_capstyle="round", zorder=7 + ) + ax.plot( + [lx, lx + 1.0], + [5.74, 5.74], + color=ORANGE, + lw=4, + solid_capstyle="round", + zorder=7, + ) + ax.text( + lx + 1.65, + 5.96, + f"cos {fmt2(cos[0])} \u2014 far apart (L0)", + color=MUTED, + fontsize=9, + family=MONO, + ha="left", + va="center", + zorder=7, + ) # tight pair = L19 - ax.plot([lx, lx + 1.0], [5.28, 5.28], color=CYAN, lw=4, solid_capstyle="round", zorder=7) - ax.plot([lx, lx + 1.0], [5.14, 5.14], color=ORANGE, lw=4, solid_capstyle="round", zorder=7) - ax.text(lx + 1.65, 5.21, f"cos {fmt2(cos[19])} \u2014 almost touching (L19)", - color=MUTED, fontsize=9, family=MONO, ha="left", va="center", zorder=7) + ax.plot( + [lx, lx + 1.0], [5.28, 5.28], color=CYAN, lw=4, solid_capstyle="round", zorder=7 + ) + ax.plot( + [lx, lx + 1.0], + [5.14, 5.14], + color=ORANGE, + lw=4, + solid_capstyle="round", + zorder=7, + ) + ax.text( + lx + 1.65, + 5.21, + f"cos {fmt2(cos[19])} \u2014 almost touching (L19)", + color=MUTED, + fontsize=9, + family=MONO, + ha="left", + va="center", + zorder=7, + ) # ---- footer -------------------------------------------------------------- - ax.text(-1.9, -6.05, - "Across the same 12 questions the image carrier matches gold answers as often as the text carrier " - f"\u2014 image EM {d['image_em'] * 100:.1f}% vs text EM {d['text_em'] * 100:.1f}%.", - color=MUTED, fontsize=9.5, family=SANS, ha="left", zorder=7) - ax.text(-1.9, -6.58, - "Data: results/qwen-carrier-convergence-n12/summary.json (29 layers \u00b7 12 SQuAD questions) " - "+ results/qwen-logit-lens-q3/logit_lens.json \u00b7 Qwen2.5-VL-7B-Instruct \u00b7 agent r2-metro", - color=MUTED, alpha=0.7, fontsize=8.5, family=MONO, ha="left", zorder=7) + ax.text( + -1.9, + -6.05, + "Across the same 12 questions the image carrier matches gold answers as often as the text carrier " + f"\u2014 image EM {d['image_em'] * 100:.1f}% vs text EM {d['text_em'] * 100:.1f}%.", + color=MUTED, + fontsize=9.5, + family=SANS, + ha="left", + zorder=7, + ) + ax.text( + -1.9, + -6.58, + "Data: results/qwen-carrier-convergence-n12/summary.json (29 layers \u00b7 12 SQuAD questions) " + "+ results/qwen-logit-lens-q3/logit_lens.json \u00b7 Qwen2.5-VL-7B-Instruct \u00b7 agent r2-metro", + color=MUTED, + alpha=0.7, + fontsize=8.5, + family=MONO, + ha="left", + zorder=7, + ) os.makedirs(OUT_DIR, exist_ok=True) fig.savefig(OUT_PNG, dpi=100, facecolor=BG) diff --git a/packages/snapcompact/research/snapcompact_tensor_heatmap.py b/packages/snapcompact/research/snapcompact_tensor_heatmap.py index 0cc441e30..840b94507 100644 --- a/packages/snapcompact/research/snapcompact_tensor_heatmap.py +++ b/packages/snapcompact/research/snapcompact_tensor_heatmap.py @@ -29,7 +29,11 @@ sys.path.insert(0, str(HERE)) import squad # noqa: E402 from bdf import capacity, render # noqa: E402 from run import CACHE, FONTS, load_prompt # noqa: E402 -from snapcompact_blackbox_occlusion import mask_cells, random_span, sample_answer_questions # noqa: E402 +from snapcompact_blackbox_occlusion import ( + mask_cells, + random_span, + sample_answer_questions, +) # noqa: E402 DEFAULT_MODEL_DIR = ( "/home/can/.cache/huggingface/hub/models--PaddlePaddle--PaddleOCR-VL/" @@ -49,10 +53,16 @@ PALETTE = { } -def ui_font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: +def ui_font( + size: int, bold: bool = False +) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: candidates = [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ] for path in candidates: if path and Path(path).exists(): @@ -96,9 +106,18 @@ def normalize(arr: np.ndarray, scale: float | None = None) -> tuple[np.ndarray, return np.clip(arr / scale, 0, 1), scale -def draw_heatmap(draw: ImageDraw.ImageDraw, arr: np.ndarray, box: tuple[int, int, int, int], title: str, subtitle: str, color: tuple[int, int, int]) -> None: +def draw_heatmap( + draw: ImageDraw.ImageDraw, + arr: np.ndarray, + box: tuple[int, int, int, int], + title: str, + subtitle: str, + color: tuple[int, int, int], +) -> None: x0, y0, x1, y1 = box - draw.rounded_rectangle(box, radius=22, fill=PALETTE["panel"], outline=(31, 42, 50), width=1) + draw.rounded_rectangle( + box, radius=22, fill=PALETTE["panel"], outline=(31, 42, 50), width=1 + ) draw.text((x0 + 24, y0 + 18), title, fill=color, font=ui_font(26, True)) draw.text((x0 + 24, y0 + 50), subtitle, fill=PALETTE["muted"], font=ui_font(15)) hx0, hy0, hx1, hy1 = x0 + 58, y0 + 84, x1 - 28, y1 - 44 @@ -116,10 +135,23 @@ def draw_heatmap(draw: ImageDraw.ImageDraw, arr: np.ndarray, box: tuple[int, int y = round(hy0 + (r + 0.5) * ch) draw.text((x0 + 18, y - 8), str(r), fill=PALETTE["muted"], font=ui_font(12)) draw.text((x0 + 16, hy0 - 4), "layer", fill=PALETTE["muted"], font=ui_font(12)) - draw.text((hx0, y1 - 31), "image token sequence →", fill=PALETTE["muted"], font=ui_font(13)) + draw.text( + (hx0, y1 - 31), + "image token sequence →", + fill=PALETTE["muted"], + font=ui_font(13), + ) -def crop_with_box(img: Image.Image, start: int, end: int, cols: int, adv: int, pitch: int, pad_cells: int = 34) -> Image.Image: +def crop_with_box( + img: Image.Image, + start: int, + end: int, + cols: int, + adv: int, + pitch: int, + pad_cells: int = 34, +) -> Image.Image: row0 = max(0, start // cols - 5) row1 = min(img.height // pitch, end // cols + 6) col0 = max(0, start % cols - pad_cells) @@ -137,34 +169,65 @@ def crop_with_box(img: Image.Image, start: int, end: int, cols: int, adv: int, p return crop -def paste_fit(canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int]) -> None: +def paste_fit( + canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int] +) -> None: x0, y0, x1, y1 = box scale = min((x1 - x0) / img.width, (y1 - y0) / img.height) - resized = img.resize((max(1, round(img.width * scale)), max(1, round(img.height * scale))), Image.Resampling.NEAREST) - canvas.paste(resized, (x0 + (x1 - x0 - resized.width) // 2, y0 + (y1 - y0 - resized.height) // 2)) + resized = img.resize( + (max(1, round(img.width * scale)), max(1, round(img.height * scale))), + Image.Resampling.NEAREST, + ) + canvas.paste( + resized, + (x0 + (x1 - x0 - resized.width) // 2, y0 + (y1 - y0 - resized.height) // 2), + ) def make_prompt(q: str, cols: int, rows: int) -> str: - return load_prompt("qa-image.md").format(cols=cols, rows=rows) + f"\n\nQuestion: {q}\nAnswer with only the shortest extractive answer." + return ( + load_prompt("qa-image.md").format(cols=cols, rows=rows) + + f"\n\nQuestion: {q}\nAnswer with only the shortest extractive answer." + ) def to_device(batch: dict[str, Any], device: Any) -> dict[str, Any]: return {k: (v.to(device) if hasattr(v, "to") else v) for k, v in batch.items()} -def hidden_token_matrix(model: Any, processor: Any, image: Image.Image, prompt_text: str, device: Any) -> tuple[list[np.ndarray], list[int], dict[str, Any]]: +def hidden_token_matrix( + model: Any, processor: Any, image: Image.Image, prompt_text: str, device: Any +) -> tuple[list[np.ndarray], list[int], dict[str, Any]]: import torch - messages = [{"role": "user", "content": [{"type": "image", "image": image}, {"type": "text", "text": prompt_text}]}] - templated = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) + messages = [ + { + "role": "user", + "content": [ + {"type": "image", "image": image}, + {"type": "text", "text": prompt_text}, + ], + } + ] + templated = processor.apply_chat_template( + messages, tokenize=False, add_generation_prompt=True + ) batch = processor(images=image, text=templated, return_tensors="pt") image_token_id = processor.tokenizer.convert_tokens_to_ids(processor.image_token) ids = batch["input_ids"][0].tolist() - image_positions = [i for i, token_id in enumerate(ids) if token_id == image_token_id] - meta = {k: (v.tolist() if hasattr(v, "tolist") else v) for k, v in batch.items() if k in ("image_grid_thw",)} + image_positions = [ + i for i, token_id in enumerate(ids) if token_id == image_token_id + ] + meta = { + k: (v.tolist() if hasattr(v, "tolist") else v) + for k, v in batch.items() + if k in ("image_grid_thw",) + } batch = to_device(batch, device) with torch.no_grad(): - out = model(**batch, output_hidden_states=True, output_attentions=False, use_cache=False) + out = model( + **batch, output_hidden_states=True, output_attentions=False, use_cache=False + ) matrices: list[np.ndarray] = [] for hidden in out.hidden_states: token_hidden = hidden[0, image_positions, :].float().detach().cpu().numpy() @@ -194,24 +257,70 @@ def render_tensor_card( gd = ImageDraw.Draw(glow) gd.ellipse((-260, -180, 850, 640), fill=(255, 83, 62, 30)) gd.ellipse((1080, 110, 2240, 1320), fill=(77, 218, 255, 30)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(80))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(80)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) - draw.text((58, 38), "SNAPCOMPACT WHITEBOX", fill=PALETTE["amber"], font=ui_font(22, True)) - draw.text((58, 76), "The hidden-state scar of a missing answer", fill=PALETTE["ink"], font=ui_font(58, True)) - draw.text((60, 148), "Each pixel below is a decoder layer × image-token bin. Bright = larger ||hidden(original) − hidden(masked)||.", fill=PALETTE["muted"], font=ui_font(24)) + draw.text( + (58, 38), "SNAPCOMPACT WHITEBOX", fill=PALETTE["amber"], font=ui_font(22, True) + ) + draw.text( + (58, 76), + "The hidden-state scar of a missing answer", + fill=PALETTE["ink"], + font=ui_font(58, True), + ) + draw.text( + (60, 148), + "Each pixel below is a decoder layer × image-token bin. Bright = larger ||hidden(original) − hidden(masked)||.", + fill=PALETTE["muted"], + font=ui_font(24), + ) # Left evidence panel. - draw.rounded_rectangle((58, 205, 700, 1098), radius=28, fill=PALETTE["panel"], outline=(31, 42, 50), width=1) - draw.text((90, 236), "the visual intervention", fill=PALETTE["ink"], font=ui_font(30, True)) - draw.text((90, 274), "same prompt, same bitmap; only answer cells blanked", fill=PALETTE["muted"], font=ui_font(17)) - crop = crop_with_box(base_img, record["answer_start"], record["answer_end"], cols, adv, pitch) - masked_crop = crop_with_box(answer_img, record["answer_start"], record["answer_end"], cols, adv, pitch) + draw.rounded_rectangle( + (58, 205, 700, 1098), + radius=28, + fill=PALETTE["panel"], + outline=(31, 42, 50), + width=1, + ) + draw.text( + (90, 236), + "the visual intervention", + fill=PALETTE["ink"], + font=ui_font(30, True), + ) + draw.text( + (90, 274), + "same prompt, same bitmap; only answer cells blanked", + fill=PALETTE["muted"], + font=ui_font(17), + ) + crop = crop_with_box( + base_img, record["answer_start"], record["answer_end"], cols, adv, pitch + ) + masked_crop = crop_with_box( + answer_img, record["answer_start"], record["answer_end"], cols, adv, pitch + ) draw.text((90, 326), "ORIGINAL", fill=PALETTE["cyan"], font=ui_font(16, True)) - draw.rounded_rectangle((90, 352, 668, 528), radius=14, fill=(244, 242, 230), outline=PALETTE["cyan"], width=3) + draw.rounded_rectangle( + (90, 352, 668, 528), + radius=14, + fill=(244, 242, 230), + outline=PALETTE["cyan"], + width=3, + ) paste_fit(canvas, crop, (108, 368, 650, 512)) draw.text((90, 568), "ANSWER ERASED", fill=PALETTE["red"], font=ui_font(16, True)) - draw.rounded_rectangle((90, 594, 668, 770), radius=14, fill=(244, 242, 230), outline=PALETTE["red"], width=3) + draw.rounded_rectangle( + (90, 594, 668, 770), + radius=14, + fill=(244, 242, 230), + outline=PALETTE["red"], + width=3, + ) paste_fit(canvas, masked_crop, (108, 610, 650, 754)) question = record["q"] if len(question) > 72: @@ -219,12 +328,43 @@ def render_tensor_card( draw.text((90, 828), "question", fill=PALETTE["muted"], font=ui_font(16, True)) draw.text((90, 856), question, fill=PALETTE["ink"], font=ui_font(21)) draw.text((90, 914), "gold answer", fill=PALETTE["muted"], font=ui_font(16, True)) - draw.text((90, 942), str(record["answer_text"]), fill=PALETTE["amber"], font=ui_font(32, True)) - draw.text((90, 1014), f"{summary['layers']} hidden layers × {summary['image_tokens']} image tokens", fill=PALETTE["muted"], font=ui_font(18)) + draw.text( + (90, 942), + str(record["answer_text"]), + fill=PALETTE["amber"], + font=ui_font(32, True), + ) + draw.text( + (90, 1014), + f"{summary['layers']} hidden layers × {summary['image_tokens']} image tokens", + fill=PALETTE["muted"], + font=ui_font(18), + ) - draw_heatmap(draw, answer_heat, (742, 205, 1818, 488), "gold answer mask", "activation delta when the true answer is blanked", PALETTE["red"]) - draw_heatmap(draw, random_heat, (742, 520, 1818, 803), "random equal-size mask", "control: blank the same number of glyph cells elsewhere", PALETTE["green"]) - draw_heatmap(draw, ratio_heat, (742, 835, 1818, 1098), "answer / random ratio", "bright bands mark layers/tokens more sensitive to the answer region", PALETTE["amber"]) + draw_heatmap( + draw, + answer_heat, + (742, 205, 1818, 488), + "gold answer mask", + "activation delta when the true answer is blanked", + PALETTE["red"], + ) + draw_heatmap( + draw, + random_heat, + (742, 520, 1818, 803), + "random equal-size mask", + "control: blank the same number of glyph cells elsewhere", + PALETTE["green"], + ) + draw_heatmap( + draw, + ratio_heat, + (742, 835, 1818, 1098), + "answer / random ratio", + "bright bands mark layers/tokens more sensitive to the answer region", + PALETTE["amber"], + ) # Color scale. for i in range(220): @@ -273,37 +413,67 @@ def main() -> None: fill = (255, 255, 255) if args.variant not in ("dark", "dark-sent") else (0, 0, 0) span_len = max(1, q["answer_end"] - q["answer_start"]) rng = random.Random(args.seed * 101 + args.question_index) - rand_start, rand_end = random_span(rng, len(chunk), span_len, q["answer_start"], q["answer_end"]) - answer_img = mask_cells(base_img, q["answer_start"], q["answer_end"], cols, cfg.adv, cfg.pitch, fill) - random_img = mask_cells(base_img, rand_start, rand_end, cols, cfg.adv, cfg.pitch, fill) + rand_start, rand_end = random_span( + rng, len(chunk), span_len, q["answer_start"], q["answer_end"] + ) + answer_img = mask_cells( + base_img, q["answer_start"], q["answer_end"], cols, cfg.adv, cfg.pitch, fill + ) + random_img = mask_cells( + base_img, rand_start, rand_end, cols, cfg.adv, cfg.pitch, fill + ) base_img.save(img_dir / "original.png") answer_img.save(img_dir / "answer-mask.png") random_img.save(img_dir / "random-mask.png") print(f"loading {args.model_dir}", flush=True) - processor = AutoProcessor.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False) + processor = AutoProcessor.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False + ) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") dtype = torch.bfloat16 if device.type == "cuda" else torch.float32 - model = AutoModel.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, dtype=dtype).to(device).eval() + model = ( + AutoModel.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True, dtype=dtype + ) + .to(device) + .eval() + ) prompt = make_prompt(q["q"], cols, rows) - original, positions, meta = hidden_token_matrix(model, processor, base_img, prompt, device) - answer, answer_positions, _ = hidden_token_matrix(model, processor, answer_img, prompt, device) - random_mask, random_positions, _ = hidden_token_matrix(model, processor, random_img, prompt, device) + original, positions, meta = hidden_token_matrix( + model, processor, base_img, prompt, device + ) + answer, answer_positions, _ = hidden_token_matrix( + model, processor, answer_img, prompt, device + ) + random_mask, random_positions, _ = hidden_token_matrix( + model, processor, random_img, prompt, device + ) if positions != answer_positions or positions != random_positions: raise SystemExit("image token positions changed across variants") - answer_delta = np.stack([np.linalg.norm(a - b, axis=1) for a, b in zip(original, answer)], axis=0) - random_delta = np.stack([np.linalg.norm(a - b, axis=1) for a, b in zip(original, random_mask)], axis=0) + answer_delta = np.stack( + [np.linalg.norm(a - b, axis=1) for a, b in zip(original, answer)], axis=0 + ) + random_delta = np.stack( + [np.linalg.norm(a - b, axis=1) for a, b in zip(original, random_mask)], axis=0 + ) ratio = answer_delta / np.maximum(random_delta, 1e-6) answer_binned = downsample_cols(answer_delta, args.bins) random_binned = downsample_cols(random_delta, args.bins) ratio_binned = downsample_cols(ratio, args.bins) - common_scale = float(np.quantile(np.concatenate([answer_binned.ravel(), random_binned.ravel()]), 0.98)) + common_scale = float( + np.quantile( + np.concatenate([answer_binned.ravel(), random_binned.ravel()]), 0.98 + ) + ) answer_norm, _ = normalize(answer_binned, common_scale) random_norm, _ = normalize(random_binned, common_scale) - ratio_norm, ratio_scale = normalize(ratio_binned, float(np.quantile(ratio_binned, 0.98))) + ratio_norm, ratio_scale = normalize( + ratio_binned, float(np.quantile(ratio_binned, 0.98)) + ) record = { "q": q["q"], @@ -325,7 +495,9 @@ def main() -> None: "processor_meta": meta, "answer_delta_mean": float(answer_delta.mean()), "random_delta_mean": float(random_delta.mean()), - "answer_over_random_delta": float(answer_delta.mean() / max(random_delta.mean(), 1e-6)), + "answer_over_random_delta": float( + answer_delta.mean() / max(random_delta.mean(), 1e-6) + ), "common_delta_scale_p98": common_scale, "ratio_scale_p98": ratio_scale, "max_ratio_layer": int(np.argmax(ratio.mean(axis=1))), @@ -345,7 +517,19 @@ def main() -> None: ratio_norm=ratio_norm, ) (out_dir / "summary.json").write_text(json.dumps(summary, indent=1)) - render_tensor_card(out_dir / "tensor-heatmap.png", answer_norm, random_norm, ratio_norm, base_img, answer_img, record, cols, cfg.adv, cfg.pitch, summary) + render_tensor_card( + out_dir / "tensor-heatmap.png", + answer_norm, + random_norm, + ratio_norm, + base_img, + answer_img, + record, + cols, + cfg.adv, + cfg.pitch, + summary, + ) print(json.dumps(summary, indent=1)) print(f"results -> {out_dir}") diff --git a/packages/snapcompact/research/snapcompact_text_image_3d_viz.py b/packages/snapcompact/research/snapcompact_text_image_3d_viz.py index 12723a4f4..5c26ae2c1 100644 --- a/packages/snapcompact/research/snapcompact_text_image_3d_viz.py +++ b/packages/snapcompact/research/snapcompact_text_image_3d_viz.py @@ -11,6 +11,7 @@ import json from pathlib import Path import matplotlib + matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np @@ -30,8 +31,12 @@ GREEN = (148, 255, 117) def font(size: int, bold: bool = False): for path in [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ]: if Path(path).exists(): return ImageFont.truetype(path, size) @@ -39,7 +44,10 @@ def font(size: int, bold: bool = False): def mono(size: int): - for path in ["/System/Library/Fonts/Monaco.ttf", "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf"]: + for path in [ + "/System/Library/Fonts/Monaco.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf", + ]: if Path(path).exists(): return ImageFont.truetype(path, size) return ImageFont.load_default() @@ -74,12 +82,45 @@ def render_surface(z: np.ndarray, answer_bins: list[int]) -> Image.Image: x = np.arange(z.shape[1]) X, Y = np.meshgrid(x, y) cmap = plt.colormaps["turbo"] - ax.plot_surface(X, Y, z, facecolors=cmap(z), linewidth=0, antialiased=True, shade=False, alpha=0.98) - ax.contour(X, Y, z, zdir="z", offset=-0.08, levels=12, cmap=cmap, linewidths=0.9, alpha=0.75) + ax.plot_surface( + X, + Y, + z, + facecolors=cmap(z), + linewidth=0, + antialiased=True, + shade=False, + alpha=0.98, + ) + ax.contour( + X, + Y, + z, + zdir="z", + offset=-0.08, + levels=12, + cmap=cmap, + linewidths=0.9, + alpha=0.75, + ) for b in answer_bins: if 0 <= b < z.shape[1]: - ax.plot([b, b], [0, z.shape[0] - 1], [1.08, 1.08], color="#ff7048", linewidth=2.6, alpha=0.78) - ax.plot([b, b], [0, z.shape[0] - 1], [-0.06, -0.06], color="#ff7048", linewidth=1.6, alpha=0.55) + ax.plot( + [b, b], + [0, z.shape[0] - 1], + [1.08, 1.08], + color="#ff7048", + linewidth=2.6, + alpha=0.78, + ) + ax.plot( + [b, b], + [0, z.shape[0] - 1], + [-0.06, -0.06], + color="#ff7048", + linewidth=1.6, + alpha=0.55, + ) ax.view_init(elev=32, azim=-58) ax.set_box_aspect((3.2, 0.8, 0.72)) ax.set_zlim(-0.08, 1.08) @@ -91,7 +132,14 @@ def render_surface(z: np.ndarray, answer_bins: list[int]) -> Image.Image: for axis in (ax.xaxis, ax.yaxis, ax.zaxis): axis.pane.set_facecolor((0.02, 0.025, 0.035, 0.0)) axis._axinfo["grid"]["color"] = (0.35, 0.45, 0.50, 0.18) - ax.set_title("text-answer vector ↔ image-token field", color="#efeede", fontsize=24, fontweight="bold", loc="left", pad=18) + ax.set_title( + "text-answer vector ↔ image-token field", + color="#efeede", + fontsize=24, + fontweight="bold", + loc="left", + pad=18, + ) tmp = HERE / "results" / ".text-image-3d-panel.png" fig.savefig(tmp, facecolor=fig.get_facecolor(), transparent=False) plt.close(fig) @@ -100,7 +148,9 @@ def render_surface(z: np.ndarray, answer_bins: list[int]) -> Image.Image: return img -def crop_answer(img: Image.Image, q: dict, cols: int, adv: int = 8, pitch: int = 13) -> Image.Image: +def crop_answer( + img: Image.Image, q: dict, cols: int, adv: int = 8, pitch: int = 13 +) -> Image.Image: start = q["answer_start"] end = q["answer_end"] row0 = max(0, start // cols - 5) @@ -117,26 +167,54 @@ def crop_answer(img: Image.Image, q: dict, cols: int, adv: int = 8, pitch: int = return crop -def paste_fit(canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int]) -> None: +def paste_fit( + canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int] +) -> None: x0, y0, x1, y1 = box scale = min((x1 - x0) / img.width, (y1 - y0) / img.height) - resized = img.resize((max(1, round(img.width * scale)), max(1, round(img.height * scale))), Image.Resampling.NEAREST) - canvas.paste(resized, (x0 + (x1 - x0 - resized.width) // 2, y0 + (y1 - y0 - resized.height) // 2)) + resized = img.resize( + (max(1, round(img.width * scale)), max(1, round(img.height * scale))), + Image.Resampling.NEAREST, + ) + canvas.paste( + resized, + (x0 + (x1 - x0 - resized.width) // 2, y0 + (y1 - y0 - resized.height) // 2), + ) def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--result-dir", default=str(HERE / "results" / "text-image-compare-paddleocr-q7")) - ap.add_argument("--out", default=str(HERE / "results" / "text-image-compare-paddleocr-q7" / "text-vs-image-3d.png")) + ap.add_argument( + "--result-dir", + default=str(HERE / "results" / "text-image-compare-paddleocr-q7"), + ) + ap.add_argument( + "--out", + default=str( + HERE + / "results" + / "text-image-compare-paddleocr-q7" + / "text-vs-image-3d.png" + ), + ) ap.add_argument("--bins", type=int, default=150) args = ap.parse_args() result_dir = Path(args.result_dir) summary = json.loads((result_dir / "summary.json").read_text()) data = np.load(result_dir / "text_image_compare.npz") - raw = data["text_answer_to_image_excess"] if "text_answer_to_image_excess" in data else data["text_answer_to_image_cosine"] + raw = ( + data["text_answer_to_image_excess"] + if "text_answer_to_image_excess" in data + else data["text_answer_to_image_cosine"] + ) z = normalize(downsample(raw, args.bins)) token_count = summary["image_tokens"] - answer_bins = sorted({round(idx / max(1, token_count - 1) * (args.bins - 1)) for idx in summary["image_answer_token_indices"]}) + answer_bins = sorted( + { + round(idx / max(1, token_count - 1) * (args.bins - 1)) + for idx in summary["image_answer_token_indices"] + } + ) panel = render_surface(z, answer_bins) w, h = 2200, 1320 @@ -148,15 +226,29 @@ def main() -> None: gd = ImageDraw.Draw(glow) gd.ellipse((-260, -220, 860, 680), fill=(75, 220, 255, 30)) gd.ellipse((1160, 80, 2440, 1320), fill=(255, 112, 72, 28)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(84))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(84)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) q = summary["question"] draw.text((64, 42), "TEXT ↔ IMAGE WHITEBOX", fill=AMBER, font=font(24, True)) - draw.text((64, 84), "Same input, different carrier, shared hidden space", fill=INK, font=font(61, True)) - draw.text((66, 166), "For every decoder layer, compare the raw-text answer state against all bitmap image-token states. Peaks = image regions whose hidden state becomes text-like.", fill=MUTED, font=font(24)) + draw.text( + (64, 84), + "Same input, different carrier, shared hidden space", + fill=INK, + font=font(61, True), + ) + draw.text( + (66, 166), + "For every decoder layer, compare the raw-text answer state against all bitmap image-token states. Peaks = image regions whose hidden state becomes text-like.", + fill=MUTED, + font=font(24), + ) - draw.rounded_rectangle((64, 238, 618, 1234), radius=30, fill=PANEL, outline=(35, 49, 59), width=1) + draw.rounded_rectangle( + (64, 238, 618, 1234), radius=30, fill=PANEL, outline=(35, 49, 59), width=1 + ) draw.text((96, 270), "two carriers", fill=INK, font=font(34, True)) draw.text((96, 312), "same chunk + same question", fill=MUTED, font=font(18)) draw.text((96, 366), "RAW TEXT", fill=CYAN, font=font(18, True)) @@ -170,28 +262,70 @@ def main() -> None: draw.text((96, y), "Gold answer token span:", fill=MUTED, font=font(17, True)) y += 32 answer_text = str(q["answer_text"]) - draw.rounded_rectangle((96, y, 108 + max(72, len(answer_text) * 24), y + 40), radius=7, fill=AMBER) + draw.rounded_rectangle( + (96, y, 108 + max(72, len(answer_text) * 24), y + 40), radius=7, fill=AMBER + ) draw.text((108, y + 7), answer_text, fill=(5, 7, 10), font=mono(22)) y += 66 draw.text((96, y), "The raw-text run receives the same", fill=INK, font=font(18)) - draw.text((96, y + 28), "SQuAD passage as ordinary tokens;", fill=INK, font=font(18)) - draw.text((96, y + 56), "the image run receives the passage", fill=INK, font=font(18)) + draw.text( + (96, y + 28), "SQuAD passage as ordinary tokens;", fill=INK, font=font(18) + ) + draw.text( + (96, y + 56), "the image run receives the passage", fill=INK, font=font(18) + ) draw.text((96, y + 84), "only through the bitmap carrier.", fill=INK, font=font(18)) - draw.text((96, 674), f"text reference: {summary['text_reference_tokens']} tokens", fill=MUTED, font=font(18)) - draw.text((96, 704), f"answer span: {summary['text_answer_tokens']} text tokens", fill=MUTED, font=font(18)) + draw.text( + (96, 674), + f"text reference: {summary['text_reference_tokens']} tokens", + fill=MUTED, + font=font(18), + ) + draw.text( + (96, 704), + f"answer span: {summary['text_answer_tokens']} text tokens", + fill=MUTED, + font=font(18), + ) draw.text((96, 774), "SNAPCOMPACT IMAGE", fill=ORANGE, font=font(18, True)) img = Image.open(result_dir / "images" / "image-carrier.png").convert("RGB") crop = crop_answer(img, q, summary["geometry"]["cols"]) - draw.rounded_rectangle((96, 812, 586, 1052), radius=16, fill=(244, 242, 230), outline=ORANGE, width=3) + draw.rounded_rectangle( + (96, 812, 586, 1052), radius=16, fill=(244, 242, 230), outline=ORANGE, width=3 + ) paste_fit(canvas, crop, (112, 828, 570, 1036)) - draw.text((96, 1092), f"image field: {summary['image_tokens']} tokens ({summary['image_grid']}×{summary['image_grid']})", fill=MUTED, font=font(18)) - draw.text((96, 1138), f"peak alignment: {summary['answer_region_cosine_max']:.3f} @ layer {summary['answer_region_cosine_argmax']}", fill=AMBER, font=font(22, True)) - draw.text((96, 1172), f"final alignment: {summary['answer_region_cosine_final']:.3f}", fill=MUTED, font=font(19)) + draw.text( + (96, 1092), + f"image field: {summary['image_tokens']} tokens ({summary['image_grid']}×{summary['image_grid']})", + fill=MUTED, + font=font(18), + ) + draw.text( + (96, 1138), + f"peak alignment: {summary['answer_region_cosine_max']:.3f} @ layer {summary['answer_region_cosine_argmax']}", + fill=AMBER, + font=font(22, True), + ) + draw.text( + (96, 1172), + f"final alignment: {summary['answer_region_cosine_final']:.3f}", + fill=MUTED, + font=font(19), + ) - draw.rounded_rectangle((650, 238, 2134, 1234), radius=30, fill=PANEL, outline=(35, 49, 59), width=1) - draw.text((686, 270), "3D cross-carrier resonance terrain", fill=INK, font=font(36, True)) - draw.text((686, 314), "z-axis = excess cosine after subtracting each layer's median image-token similarity; orange rails mark the bitmap answer region", fill=MUTED, font=font(20)) + draw.rounded_rectangle( + (650, 238, 2134, 1234), radius=30, fill=PANEL, outline=(35, 49, 59), width=1 + ) + draw.text( + (686, 270), "3D cross-carrier resonance terrain", fill=INK, font=font(36, True) + ) + draw.text( + (686, 314), + "z-axis = excess cosine after subtracting each layer's median image-token similarity; orange rails mark the bitmap answer region", + fill=MUTED, + font=font(20), + ) panel = panel.resize((1408, 794), Image.Resampling.LANCZOS) canvas.paste(panel, (692, 378)) cmap = plt.colormaps["turbo"] diff --git a/packages/snapcompact/research/snapcompact_text_image_compare.py b/packages/snapcompact/research/snapcompact_text_image_compare.py index 000e46d81..0b5a62eea 100644 --- a/packages/snapcompact/research/snapcompact_text_image_compare.py +++ b/packages/snapcompact/research/snapcompact_text_image_compare.py @@ -53,11 +53,17 @@ PALETTE = { } -def ui_font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: +def ui_font( + size: int, bold: bool = False +) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: candidates = [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", "/System/Library/Fonts/Monaco.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ] for path in candidates: if path and Path(path).exists(): @@ -99,7 +105,9 @@ def cosine(a: np.ndarray, b: np.ndarray) -> np.ndarray: return (a * b).sum(axis=-1) / np.maximum((a_norm * b_norm).squeeze(-1), 1e-6) -def normalize_heat(arr: np.ndarray, lo: float | None = None, hi: float | None = None) -> tuple[np.ndarray, float, float]: +def normalize_heat( + arr: np.ndarray, lo: float | None = None, hi: float | None = None +) -> tuple[np.ndarray, float, float]: if lo is None: lo = float(np.quantile(arr, 0.03)) if hi is None: @@ -110,15 +118,23 @@ def normalize_heat(arr: np.ndarray, lo: float | None = None, hi: float | None = def apply_template(processor: Any, content: list[dict[str, Any]]) -> str: - return processor.apply_chat_template([{"role": "user", "content": content}], tokenize=False, add_generation_prompt=True) + return processor.apply_chat_template( + [{"role": "user", "content": content}], + tokenize=False, + add_generation_prompt=True, + ) -def text_spans(processor: Any, templated: str, chunk: str, answer_start: int, answer_end: int) -> dict[str, int]: +def text_spans( + processor: Any, templated: str, chunk: str, answer_start: int, answer_end: int +) -> dict[str, int]: tokenizer = processor.tokenizer chunk_at = templated.index(chunk) prefix = templated[:chunk_at] + def n_tokens(s: str) -> int: return len(tokenizer(s, add_special_tokens=False)["input_ids"]) + ref_start = n_tokens(prefix) ref_end = n_tokens(prefix + chunk) answer_tok_start = n_tokens(prefix + chunk[:answer_start]) @@ -135,7 +151,15 @@ def to_device(batch: dict[str, Any], device: Any) -> dict[str, Any]: return {k: (v.to(device) if hasattr(v, "to") else v) for k, v in batch.items()} -def run_text(model: Any, processor: Any, text_prompt: str, chunk: str, answer_start: int, answer_end: int, device: Any) -> tuple[list[np.ndarray], dict[str, int], str]: +def run_text( + model: Any, + processor: Any, + text_prompt: str, + chunk: str, + answer_start: int, + answer_end: int, + device: Any, +) -> tuple[list[np.ndarray], dict[str, int], str]: import torch templated = apply_template(processor, [{"type": "text", "text": text_prompt}]) @@ -143,27 +167,59 @@ def run_text(model: Any, processor: Any, text_prompt: str, chunk: str, answer_st batch = processor(text=templated, return_tensors="pt") batch = to_device(batch, device) with torch.no_grad(): - out = model(**batch, output_hidden_states=True, output_attentions=False, use_cache=False) - layers = [h[0].float().detach().cpu().numpy().astype(np.float32, copy=False) for h in out.hidden_states] + out = model( + **batch, output_hidden_states=True, output_attentions=False, use_cache=False + ) + layers = [ + h[0].float().detach().cpu().numpy().astype(np.float32, copy=False) + for h in out.hidden_states + ] return layers, spans, templated -def run_image(model: Any, processor: Any, img: Image.Image, img_prompt: str, device: Any) -> tuple[list[np.ndarray], list[int], dict[str, Any], str]: +def run_image( + model: Any, processor: Any, img: Image.Image, img_prompt: str, device: Any +) -> tuple[list[np.ndarray], list[int], dict[str, Any], str]: import torch - templated = apply_template(processor, [{"type": "image", "image": img}, {"type": "text", "text": img_prompt}]) + templated = apply_template( + processor, + [{"type": "image", "image": img}, {"type": "text", "text": img_prompt}], + ) batch = processor(images=img, text=templated, return_tensors="pt") image_token_id = processor.tokenizer.convert_tokens_to_ids(processor.image_token) - image_positions = [i for i, token_id in enumerate(batch["input_ids"][0].tolist()) if token_id == image_token_id] - meta = {k: (v.tolist() if hasattr(v, "tolist") else v) for k, v in batch.items() if k in ("image_grid_thw",)} + image_positions = [ + i + for i, token_id in enumerate(batch["input_ids"][0].tolist()) + if token_id == image_token_id + ] + meta = { + k: (v.tolist() if hasattr(v, "tolist") else v) + for k, v in batch.items() + if k in ("image_grid_thw",) + } batch = to_device(batch, device) with torch.no_grad(): - out = model(**batch, output_hidden_states=True, output_attentions=False, use_cache=False) - layers = [h[0].float().detach().cpu().numpy().astype(np.float32, copy=False) for h in out.hidden_states] + out = model( + **batch, output_hidden_states=True, output_attentions=False, use_cache=False + ) + layers = [ + h[0].float().detach().cpu().numpy().astype(np.float32, copy=False) + for h in out.hidden_states + ] return layers, image_positions, meta, templated -def image_answer_token_indices(answer_start: int, answer_end: int, text_cols: int, adv: int, pitch: int, image_w: int, image_h: int, image_token_count: int) -> list[int]: +def image_answer_token_indices( + answer_start: int, + answer_end: int, + text_cols: int, + adv: int, + pitch: int, + image_w: int, + image_h: int, + image_token_count: int, +) -> list[int]: grid = round(math.sqrt(image_token_count)) if grid * grid != image_token_count: return [] @@ -186,7 +242,15 @@ def image_answer_token_indices(answer_start: int, answer_end: int, text_cols: in return sorted(set(out)) -def crop_answer(img: Image.Image, start: int, end: int, cols: int, adv: int, pitch: int, pad_cells: int = 34) -> Image.Image: +def crop_answer( + img: Image.Image, + start: int, + end: int, + cols: int, + adv: int, + pitch: int, + pad_cells: int = 34, +) -> Image.Image: row0 = max(0, start // cols - 5) row1 = min(img.height // pitch, end // cols + 6) col0 = max(0, start % cols - pad_cells) @@ -199,11 +263,21 @@ def crop_answer(img: Image.Image, start: int, end: int, cols: int, adv: int, pit bx1 = min(crop.width - 1, ((end - 1) % cols - col0 + 2) * adv) by0 = max(0, (start // cols - row0) * pitch - 1) by1 = min(crop.height - 1, ((end - 1) // cols - row0 + 1) * pitch + 1) - d.rounded_rectangle((bx0, by0, bx1, by1), radius=3, outline=PALETTE["orange"], width=3) + d.rounded_rectangle( + (bx0, by0, bx1, by1), radius=3, outline=PALETTE["orange"], width=3 + ) return crop -def draw_wrapped(draw: ImageDraw.ImageDraw, xy: tuple[int, int], text: str, width_chars: int, line_height: int, fill: tuple[int, int, int], fnt: ImageFont.ImageFont) -> int: +def draw_wrapped( + draw: ImageDraw.ImageDraw, + xy: tuple[int, int], + text: str, + width_chars: int, + line_height: int, + fill: tuple[int, int, int], + fnt: ImageFont.ImageFont, +) -> int: words = text.split() lines: list[str] = [] current = "" @@ -224,11 +298,26 @@ def draw_wrapped(draw: ImageDraw.ImageDraw, xy: tuple[int, int], text: str, widt return y -def render_heat_grid(draw: ImageDraw.ImageDraw, grid: np.ndarray, box: tuple[int, int, int, int], title: str, layer: int, answer_indices: list[int], color: tuple[int, int, int]) -> None: +def render_heat_grid( + draw: ImageDraw.ImageDraw, + grid: np.ndarray, + box: tuple[int, int, int, int], + title: str, + layer: int, + answer_indices: list[int], + color: tuple[int, int, int], +) -> None: x0, y0, x1, y1 = box - draw.rounded_rectangle(box, radius=18, fill=PALETTE["panel2"], outline=(35, 49, 59), width=1) + draw.rounded_rectangle( + box, radius=18, fill=PALETTE["panel2"], outline=(35, 49, 59), width=1 + ) draw.text((x0 + 18, y0 + 14), title, fill=color, font=ui_font(22, True)) - draw.text((x0 + 18, y0 + 43), f"decoder layer {layer}", fill=PALETTE["muted"], font=ui_font(15)) + draw.text( + (x0 + 18, y0 + 43), + f"decoder layer {layer}", + fill=PALETTE["muted"], + font=ui_font(15), + ) gx0, gy0, gx1, gy1 = x0 + 26, y0 + 76, x1 - 26, y1 - 24 rows, cols = grid.shape cw = (gx1 - gx0) / cols @@ -246,10 +335,18 @@ def render_heat_grid(draw: ImageDraw.ImageDraw, grid: np.ndarray, box: tuple[int xb = round(gx0 + (c + 1) * cw) ya = round(gy0 + r * ch) yb = round(gy0 + (r + 1) * ch) - draw.rectangle((xa - 2, ya - 2, xb + 2, yb + 2), outline=PALETTE["orange"], width=2) + draw.rectangle( + (xa - 2, ya - 2, xb + 2, yb + 2), outline=PALETTE["orange"], width=2 + ) -def render_visual(out_path: Path, summary: dict[str, Any], arrays: dict[str, np.ndarray], original_img: Image.Image, chunk: str) -> None: +def render_visual( + out_path: Path, + summary: dict[str, Any], + arrays: dict[str, np.ndarray], + original_img: Image.Image, + chunk: str, +) -> None: w, h = 2100, 1260 canvas = Image.new("RGB", (w, h), PALETTE["bg"]) draw = ImageDraw.Draw(canvas) @@ -259,17 +356,42 @@ def render_visual(out_path: Path, summary: dict[str, Any], arrays: dict[str, np. gd = ImageDraw.Draw(glow) gd.ellipse((-240, -220, 860, 660), fill=(75, 220, 255, 28)) gd.ellipse((1160, 80, 2420, 1320), fill=(255, 112, 72, 26)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(80))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(80)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) q = summary["question"] - draw.text((62, 42), "SNAPCOMPACT CARRIER COMPARISON", fill=PALETTE["amber"], font=ui_font(24, True)) - draw.text((62, 82), "Same input, two internal languages", fill=PALETTE["ink"], font=ui_font(68, True)) - draw.text((64, 166), "Raw text tokens vs bitmap image tokens. Bright fields show where the text-carrier answer vector resonates with the image-carrier hidden state.", fill=PALETTE["muted"], font=ui_font(25)) + draw.text( + (62, 42), + "SNAPCOMPACT CARRIER COMPARISON", + fill=PALETTE["amber"], + font=ui_font(24, True), + ) + draw.text( + (62, 82), + "Same input, two internal languages", + fill=PALETTE["ink"], + font=ui_font(68, True), + ) + draw.text( + (64, 166), + "Raw text tokens vs bitmap image tokens. Bright fields show where the text-carrier answer vector resonates with the image-carrier hidden state.", + fill=PALETTE["muted"], + font=ui_font(25), + ) # Carrier cards. - draw.rounded_rectangle((62, 236, 620, 760), radius=28, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((94, 270), "raw text carrier", fill=PALETTE["cyan"], font=ui_font(30, True)) + draw.rounded_rectangle( + (62, 236, 620, 760), + radius=28, + fill=PALETTE["panel"], + outline=(35, 49, 59), + width=1, + ) + draw.text( + (94, 270), "raw text carrier", fill=PALETTE["cyan"], font=ui_font(30, True) + ) start = max(0, q["answer_start"] - 230) end = min(len(chunk), q["answer_end"] + 230) snippet = chunk[start:end].replace("\n", " ") @@ -279,22 +401,69 @@ def render_visual(out_path: Path, summary: dict[str, Any], arrays: dict[str, np. answer = snippet[rel_a:rel_b] after = snippet[rel_b:] tx, ty = 94, 328 - ty = draw_wrapped(draw, (tx, ty), before[-260:], 52, 22, PALETTE["ink"], mono_font(15)) - draw.rounded_rectangle((tx, ty + 2, tx + 16 * max(3, len(answer)), ty + 27), radius=5, fill=(255, 196, 68)) + ty = draw_wrapped( + draw, (tx, ty), before[-260:], 52, 22, PALETTE["ink"], mono_font(15) + ) + draw.rounded_rectangle( + (tx, ty + 2, tx + 16 * max(3, len(answer)), ty + 27), + radius=5, + fill=(255, 196, 68), + ) draw.text((tx + 4, ty + 5), answer, fill=(8, 10, 10), font=mono_font(16)) ty += 36 draw_wrapped(draw, (tx, ty), after[:260], 52, 22, PALETTE["ink"], mono_font(15)) - draw.text((94, 694), f"answer tokens: {summary['text_answer_tokens']}", fill=PALETTE["muted"], font=ui_font(18)) - draw.text((94, 724), f"reference tokens: {summary['text_reference_tokens']}", fill=PALETTE["muted"], font=ui_font(18)) + draw.text( + (94, 694), + f"answer tokens: {summary['text_answer_tokens']}", + fill=PALETTE["muted"], + font=ui_font(18), + ) + draw.text( + (94, 724), + f"reference tokens: {summary['text_reference_tokens']}", + fill=PALETTE["muted"], + font=ui_font(18), + ) - draw.rounded_rectangle((62, 792, 620, 1192), radius=28, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((94, 826), "image carrier", fill=PALETTE["orange"], font=ui_font(30, True)) - crop = crop_answer(original_img, q["answer_start"], q["answer_end"], summary["geometry"]["cols"], 8, 13) + draw.rounded_rectangle( + (62, 792, 620, 1192), + radius=28, + fill=PALETTE["panel"], + outline=(35, 49, 59), + width=1, + ) + draw.text( + (94, 826), "image carrier", fill=PALETTE["orange"], font=ui_font(30, True) + ) + crop = crop_answer( + original_img, + q["answer_start"], + q["answer_end"], + summary["geometry"]["cols"], + 8, + 13, + ) scale = min(478 / crop.width, 218 / crop.height) - crop_r = crop.resize((round(crop.width * scale), round(crop.height * scale)), Image.Resampling.NEAREST) - draw.rounded_rectangle((94, 888, 588, 1134), radius=16, fill=(244, 242, 230), outline=PALETTE["orange"], width=3) - canvas.paste(crop_r, (94 + (494 - crop_r.width) // 2, 888 + (246 - crop_r.height) // 2)) - draw.text((94, 1150), f"image tokens: {summary['image_tokens']} ({summary['image_grid']}×{summary['image_grid']})", fill=PALETTE["muted"], font=ui_font(18)) + crop_r = crop.resize( + (round(crop.width * scale), round(crop.height * scale)), + Image.Resampling.NEAREST, + ) + draw.rounded_rectangle( + (94, 888, 588, 1134), + radius=16, + fill=(244, 242, 230), + outline=PALETTE["orange"], + width=3, + ) + canvas.paste( + crop_r, (94 + (494 - crop_r.width) // 2, 888 + (246 - crop_r.height) // 2) + ) + draw.text( + (94, 1150), + f"image tokens: {summary['image_tokens']} ({summary['image_grid']}×{summary['image_grid']})", + fill=PALETTE["muted"], + font=ui_font(18), + ) # Layer grids. sim = arrays["text_answer_to_image_excess_norm"] @@ -305,12 +474,36 @@ def render_visual(out_path: Path, summary: dict[str, Any], arrays: dict[str, np. names = ["input layer", "middle layer", "peak alignment"] colors = [PALETTE["cyan"], PALETTE["purple"], PALETTE["green"]] for layer, box, name, color in zip(layers, boxes, names, colors): - render_heat_grid(draw, sim[layer].reshape(grid, grid), box, name, layer, answer_indices, color) + render_heat_grid( + draw, + sim[layer].reshape(grid, grid), + box, + name, + layer, + answer_indices, + color, + ) # Cosine bridge panel. - draw.rounded_rectangle((672, 672, 1980, 1192), radius=28, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((704, 704), "cross-carrier convergence bridge", fill=PALETTE["ink"], font=ui_font(34, True)) - draw.text((704, 744), "Cosine similarity between pooled raw-text answer states and pooled bitmap answer-region states by layer", fill=PALETTE["muted"], font=ui_font(19)) + draw.rounded_rectangle( + (672, 672, 1980, 1192), + radius=28, + fill=PALETTE["panel"], + outline=(35, 49, 59), + width=1, + ) + draw.text( + (704, 704), + "cross-carrier convergence bridge", + fill=PALETTE["ink"], + font=ui_font(34, True), + ) + draw.text( + (704, 744), + "Cosine similarity between pooled raw-text answer states and pooled bitmap answer-region states by layer", + fill=PALETTE["muted"], + font=ui_font(19), + ) x0, y0, x1, y1 = 730, 820, 1908, 1096 for i in range(5): y = y0 + round((y1 - y0) * i / 4) @@ -321,6 +514,7 @@ def render_visual(out_path: Path, summary: dict[str, Any], arrays: dict[str, np. hi = float(max(local.max(), global_mean.max())) if hi <= lo: hi = lo + 1e-6 + def pts(vals: np.ndarray) -> list[tuple[int, int]]: out = [] for i, v in enumerate(vals): @@ -328,6 +522,7 @@ def render_visual(out_path: Path, summary: dict[str, Any], arrays: dict[str, np. y = y1 - round((y1 - y0) * (float(v) - lo) / (hi - lo)) out.append((x, y)) return out + p_local = pts(local) p_global = pts(global_mean) draw.line(p_global, fill=PALETTE["muted"], width=4) @@ -335,17 +530,56 @@ def render_visual(out_path: Path, summary: dict[str, Any], arrays: dict[str, np. for x, y in p_local: draw.ellipse((x - 5, y - 5, x + 5, y + 5), fill=PALETTE["amber"]) draw.text((x0, y1 + 22), "layer 0", fill=PALETTE["muted"], font=ui_font(16)) - draw.text((x1 - 70, y1 + 22), f"layer {len(local) - 1}", fill=PALETTE["muted"], font=ui_font(16)) + draw.text( + (x1 - 70, y1 + 22), + f"layer {len(local) - 1}", + fill=PALETTE["muted"], + font=ui_font(16), + ) peak_layer = int(np.argmax(local)) - draw.rounded_rectangle((1502, 790, 1938, 900), radius=18, fill=(9, 13, 18), outline=(38, 51, 60), width=1) - draw.text((1526, 812), f"answer cosine peaks: {local[peak_layer]:.3f} @L{peak_layer}", fill=PALETTE["amber"], font=ui_font(21, True)) - draw.text((1526, 842), f"final answer cosine: {local[-1]:.3f}", fill=PALETTE["muted"], font=ui_font(18)) - draw.text((1526, 868), f"final global carrier cosine: {global_mean[-1]:.3f}", fill=PALETTE["muted"], font=ui_font(18)) - draw.rounded_rectangle((704, 1120, 1238, 1168), radius=13, fill=(9, 13, 18), outline=(38, 51, 60), width=1) + draw.rounded_rectangle( + (1502, 790, 1938, 900), + radius=18, + fill=(9, 13, 18), + outline=(38, 51, 60), + width=1, + ) + draw.text( + (1526, 812), + f"answer cosine peaks: {local[peak_layer]:.3f} @L{peak_layer}", + fill=PALETTE["amber"], + font=ui_font(21, True), + ) + draw.text( + (1526, 842), + f"final answer cosine: {local[-1]:.3f}", + fill=PALETTE["muted"], + font=ui_font(18), + ) + draw.text( + (1526, 868), + f"final global carrier cosine: {global_mean[-1]:.3f}", + fill=PALETTE["muted"], + font=ui_font(18), + ) + draw.rounded_rectangle( + (704, 1120, 1238, 1168), + radius=13, + fill=(9, 13, 18), + outline=(38, 51, 60), + width=1, + ) draw.rectangle((724, 1138, 768, 1148), fill=PALETTE["amber"]) - draw.text((784, 1129), "answer region: text vector ↔ image region", fill=PALETTE["muted"], font=ui_font(17)) + draw.text( + (784, 1129), + "answer region: text vector ↔ image region", + fill=PALETTE["muted"], + font=ui_font(17), + ) draw.rectangle((1260, 1138, 1304, 1148), fill=PALETTE["muted"]) - draw.text((1320, 1129), "global carrier means", fill=PALETTE["muted"], font=ui_font(17)) + draw.text( + (1320, 1129), "global carrier means", fill=PALETTE["muted"], font=ui_font(17) + ) out_path.parent.mkdir(parents=True, exist_ok=True) canvas.save(out_path) @@ -386,8 +620,12 @@ def main() -> None: img.save(img_dir / "image-carrier.png") print(f"loading {args.model_dir}", flush=True) - config = AutoConfig.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True) - processor = AutoProcessor.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False) + config = AutoConfig.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True + ) + processor = AutoProcessor.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False + ) target_device = torch.device("cuda" if torch.cuda.is_available() else "cpu") dtype = torch.bfloat16 if target_device.type == "cuda" else torch.float32 if getattr(config, "model_type", "") == "qwen2_5_vl": @@ -402,7 +640,16 @@ def main() -> None: ).eval() device = next(model.parameters()).device else: - model = AutoModel.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, dtype=dtype).to(target_device).eval() + model = ( + AutoModel.from_pretrained( + args.model_dir, + local_files_only=True, + trust_remote_code=True, + dtype=dtype, + ) + .to(target_device) + .eval() + ) device = target_device text_prompt = ( @@ -410,13 +657,29 @@ def main() -> None: f"{chunk}\n\nQuestion: {q['q']}\n" "Answer with only the shortest extractive answer." ) - img_prompt = load_prompt("qa-image.md").format(cols=cols, rows=rows) + f"\n\nQuestion: {q['q']}\nAnswer with only the shortest extractive answer." + img_prompt = ( + load_prompt("qa-image.md").format(cols=cols, rows=rows) + + f"\n\nQuestion: {q['q']}\nAnswer with only the shortest extractive answer." + ) - text_layers, text_pos, text_template = run_text(model, processor, text_prompt, chunk, q["answer_start"], q["answer_end"], device) - image_layers, image_positions, image_meta, image_template = run_image(model, processor, img, img_prompt, device) + text_layers, text_pos, text_template = run_text( + model, processor, text_prompt, chunk, q["answer_start"], q["answer_end"], device + ) + image_layers, image_positions, image_meta, image_template = run_image( + model, processor, img, img_prompt, device + ) image_token_count = len(image_positions) image_grid = round(math.sqrt(image_token_count)) - answer_image_indices = image_answer_token_indices(q["answer_start"], q["answer_end"], cols, cfg.adv, cfg.pitch, img.width, img.height, image_token_count) + answer_image_indices = image_answer_token_indices( + q["answer_start"], + q["answer_end"], + cols, + cfg.adv, + cfg.pitch, + img.width, + img.height, + image_token_count, + ) answer_cos = [] global_cos = [] @@ -425,14 +688,21 @@ def main() -> None: text_ref = text_h[text_pos["ref_start"] : text_pos["ref_end"]] text_ans = text_h[text_pos["answer_start"] : text_pos["answer_end"]] image_tokens = image_h[image_positions] - image_ans = image_tokens[answer_image_indices] if answer_image_indices else image_tokens + image_ans = ( + image_tokens[answer_image_indices] if answer_image_indices else image_tokens + ) text_ans_mean = text_ans.mean(axis=0) image_ans_mean = image_ans.mean(axis=0) text_ref_mean = text_ref.mean(axis=0) image_mean = image_tokens.mean(axis=0) - answer_cos.append(float(cosine(text_ans_mean[None, :], image_ans_mean[None, :])[0])) + answer_cos.append( + float(cosine(text_ans_mean[None, :], image_ans_mean[None, :])[0]) + ) global_cos.append(float(cosine(text_ref_mean[None, :], image_mean[None, :])[0])) - sims = cosine(np.repeat(text_ans_mean[None, :], image_tokens.shape[0], axis=0), image_tokens) + sims = cosine( + np.repeat(text_ans_mean[None, :], image_tokens.shape[0], axis=0), + image_tokens, + ) text_answer_to_image.append(sims.astype(np.float32, copy=False)) text_answer_to_image_arr = np.stack(text_answer_to_image, axis=0) diff --git a/packages/snapcompact/research/snapcompact_token_entry_dump.py b/packages/snapcompact/research/snapcompact_token_entry_dump.py index 5d9516192..45cfee069 100644 --- a/packages/snapcompact/research/snapcompact_token_entry_dump.py +++ b/packages/snapcompact/research/snapcompact_token_entry_dump.py @@ -64,10 +64,20 @@ def main() -> None: print(f"loading {args.model_dir}", flush=True) from transformers import AutoTokenizer - processor = AutoProcessor.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False) - model = Qwen2_5_VLForConditionalGeneration.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True, dtype=torch.bfloat16, device_map="auto").eval() + processor = AutoProcessor.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True, use_fast=False + ) + model = Qwen2_5_VLForConditionalGeneration.from_pretrained( + args.model_dir, + local_files_only=True, + trust_remote_code=True, + dtype=torch.bfloat16, + device_map="auto", + ).eval() device = next(model.parameters()).device - tokenizer = AutoTokenizer.from_pretrained(args.model_dir, local_files_only=True, trust_remote_code=True) # fast tokenizer for offsets + tokenizer = AutoTokenizer.from_pretrained( + args.model_dir, local_files_only=True, trust_remote_code=True + ) # fast tokenizer for offsets # --- Text lane: real tokenization of the snippet around the answer. snip_start = max(0, q["answer_start"] - args.context_chars) @@ -78,22 +88,35 @@ def main() -> None: answer_token_idx: list[int] = [] rel_a = q["answer_start"] - snip_start rel_b = q["answer_end"] - snip_start - for ti, (tok_id, (o0, o1)) in enumerate(zip(enc["input_ids"], enc["offset_mapping"])): + for ti, (tok_id, (o0, o1)) in enumerate( + zip(enc["input_ids"], enc["offset_mapping"]) + ): is_answer = o0 < rel_b and o1 > rel_a if is_answer: answer_token_idx.append(ti) - tokens.append({"i": ti, "id": int(tok_id), "str": tokenizer.decode([tok_id]), "answer": bool(is_answer)}) + tokens.append( + { + "i": ti, + "id": int(tok_id), + "str": tokenizer.decode([tok_id]), + "answer": bool(is_answer), + } + ) # Real embedding rows entering the decoder for the answer tokens. embed = model.get_input_embeddings() - answer_ids = torch.tensor([tokens[i]["id"] for i in answer_token_idx], device=device) + answer_ids = torch.tensor( + [tokens[i]["id"] for i in answer_token_idx], device=device + ) with torch.no_grad(): answer_embeds = embed(answer_ids).float().cpu().numpy() text_entry = [ { "id": tokens[i]["id"], "str": tokens[i]["str"], - "vector_head": [round(float(v), 4) for v in answer_embeds[k, : args.embed_dims]], + "vector_head": [ + round(float(v), 4) for v in answer_embeds[k, : args.embed_dims] + ], "norm": round(float(np.linalg.norm(answer_embeds[k])), 4), } for k, i in enumerate(answer_token_idx) @@ -101,34 +124,74 @@ def main() -> None: chunk_token_count = len(tokenizer(chunk, add_special_tokens=False)["input_ids"]) # --- Image lane: real pixel patches and visual-tower output vectors. - batch = processor(images=img, text="<|vision_start|><|image_pad|><|vision_end|>", return_tensors="pt") + batch = processor( + images=img, + text="<|vision_start|><|image_pad|><|vision_end|>", + return_tensors="pt", + ) pixel_values = batch["pixel_values"] grid_thw = batch["image_grid_thw"] merge = int(getattr(processor.image_processor, "merge_size", 2)) patch = int(getattr(processor.image_processor, "patch_size", 14)) with torch.no_grad(): - visual_out = model.model.visual(pixel_values.to(device, dtype=torch.bfloat16), grid_thw=grid_thw.to(device)).float().cpu().numpy() + visual_out = ( + model.model.visual( + pixel_values.to(device, dtype=torch.bfloat16), + grid_thw=grid_thw.to(device), + ) + .float() + .cpu() + .numpy() + ) n_tokens = visual_out.shape[0] grid = int(round(n_tokens**0.5)) - answer_img_indices = image_answer_token_indices(q["answer_start"], q["answer_end"], cols, cfg.adv, cfg.pitch, img.width, img.height, n_tokens) + answer_img_indices = image_answer_token_indices( + q["answer_start"], + q["answer_end"], + cols, + cfg.adv, + cfg.pitch, + img.width, + img.height, + n_tokens, + ) image_entry = [ { "token_index": int(idx), "grid_rc": [int(idx // grid), int(idx % grid)], - "vector_head": [round(float(v), 4) for v in visual_out[idx, : args.embed_dims]], + "vector_head": [ + round(float(v), 4) for v in visual_out[idx, : args.embed_dims] + ], "norm": round(float(np.linalg.norm(visual_out[idx])), 4), } for idx in answer_img_indices ] # A few real normalized pixel values from the first answer patch (pre-visual-tower input). patches_per_token = merge * merge - first_patch_row = answer_img_indices[0] * patches_per_token if answer_img_indices else 0 - pixel_head = [round(float(v), 4) for v in pixel_values[min(first_patch_row, pixel_values.shape[0] - 1), : args.embed_dims].tolist()] + first_patch_row = ( + answer_img_indices[0] * patches_per_token if answer_img_indices else 0 + ) + pixel_head = [ + round(float(v), 4) + for v in pixel_values[ + min(first_patch_row, pixel_values.shape[0] - 1), : args.embed_dims + ].tolist() + ] dump = { "args": vars(args), - "question": {"q": q["q"], "answer_text": q["answer_text"], "answer_start": q["answer_start"], "answer_end": q["answer_end"]}, - "geometry": {"cols": cols, "rows": rows, "image_w": img.width, "image_h": img.height}, + "question": { + "q": q["q"], + "answer_text": q["answer_text"], + "answer_start": q["answer_start"], + "answer_end": q["answer_end"], + }, + "geometry": { + "cols": cols, + "rows": rows, + "image_w": img.width, + "image_h": img.height, + }, "snippet": snippet, "snippet_rel_answer": [rel_a, rel_b], "tokens": tokens, @@ -150,7 +213,11 @@ def main() -> None: "visual_out_dim": int(visual_out.shape[1]), } (out_dir / "token_entry.json").write_text(json.dumps(dump, indent=1)) - print(json.dumps({k: v for k, v in dump.items() if k not in ("tokens", "snippet")}, indent=1)) + print( + json.dumps( + {k: v for k, v in dump.items() if k not in ("tokens", "snippet")}, indent=1 + ) + ) print(f"results -> {out_dir}") diff --git a/packages/snapcompact/research/snapcompact_token_entry_viz.py b/packages/snapcompact/research/snapcompact_token_entry_viz.py index 3feaac0f6..9b5add19a 100644 --- a/packages/snapcompact/research/snapcompact_token_entry_viz.py +++ b/packages/snapcompact/research/snapcompact_token_entry_viz.py @@ -31,8 +31,12 @@ PALETTE = { def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: for path in [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ]: if Path(path).exists(): return ImageFont.truetype(path, size) @@ -40,7 +44,10 @@ def ui_font(size: int, bold: bool = False) -> ImageFont.ImageFont: def mono_font(size: int) -> ImageFont.ImageFont: - for path in ["/System/Library/Fonts/Monaco.ttf", "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf"]: + for path in [ + "/System/Library/Fonts/Monaco.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSansMono.ttf", + ]: if Path(path).exists(): return ImageFont.truetype(path, size) return ImageFont.load_default() @@ -50,7 +57,13 @@ def vector_text(head: list[float]) -> str: return "[" + ", ".join(f"{v:+.2f}" for v in head[:6]) + ", …]" -def draw_vector_bar(draw: ImageDraw.ImageDraw, xy: tuple[int, int], head: list[float], color: tuple[int, int, int], width: int = 330) -> None: +def draw_vector_bar( + draw: ImageDraw.ImageDraw, + xy: tuple[int, int], + head: list[float], + color: tuple[int, int, int], + width: int = 330, +) -> None: x, y = xy n = len(head) bw = width // n @@ -60,16 +73,27 @@ def draw_vector_bar(draw: ImageDraw.ImageDraw, xy: tuple[int, int], head: list[f bh = round(20 * abs(v) / hi) xa = x + i * bw if v >= 0: - draw.rounded_rectangle((xa, mid - bh, xa + bw - 4, mid), radius=3, fill=color) + draw.rounded_rectangle( + (xa, mid - bh, xa + bw - 4, mid), radius=3, fill=color + ) else: - draw.rounded_rectangle((xa, mid, xa + bw - 4, mid + bh), radius=3, fill=tuple(c // 2 for c in color)) + draw.rounded_rectangle( + (xa, mid, xa + bw - 4, mid + bh), + radius=3, + fill=tuple(c // 2 for c in color), + ) draw.line((x, mid, x + width, mid), fill=PALETTE["grid"], width=1) def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--result-dir", default=str(HERE / "results" / "qwen-token-entry-q3")) - ap.add_argument("--out", default=str(HERE / "results" / "qwen-token-entry-q3" / "token-entry.png")) + ap.add_argument( + "--result-dir", default=str(HERE / "results" / "qwen-token-entry-q3") + ) + ap.add_argument( + "--out", + default=str(HERE / "results" / "qwen-token-entry-q3" / "token-entry.png"), + ) args = ap.parse_args() result_dir = Path(args.result_dir) dump = json.loads((result_dir / "token_entry.json").read_text()) @@ -84,13 +108,25 @@ def main() -> None: gd = ImageDraw.Draw(glow) gd.ellipse((-240, -200, 940, 760), fill=(75, 220, 255, 27)) gd.ellipse((1240, 540, 2460, 1480), fill=(255, 112, 72, 25)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(86))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(86)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) q = dump["question"] answer = q["answer_text"] - draw.text((64, 42), "QWEN TOKEN ENTRY — SAME WORD, TWO ENCODINGS", fill=PALETTE["amber"], font=ui_font(24, True)) - draw.text((64, 84), f"How “{answer}” gets into the model", fill=PALETTE["ink"], font=ui_font(64, True)) + draw.text( + (64, 42), + "QWEN TOKEN ENTRY — SAME WORD, TWO ENCODINGS", + fill=PALETTE["amber"], + font=ui_font(24, True), + ) + draw.text( + (64, 84), + f"How “{answer}” gets into the model", + fill=PALETTE["ink"], + font=ui_font(64, True), + ) draw.text( (66, 164), "Real values, no schematic: actual BPE ids and embedding rows on the text path; actual 28×28 pixel patches and visual-tower output vectors on the image path.", @@ -100,9 +136,21 @@ def main() -> None: # ---- TEXT LANE ---- lane = (64, 238, 2136, 700) - draw.rounded_rectangle(lane, radius=28, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((96, 262), "text carrier — BPE tokens", fill=PALETTE["cyan"], font=ui_font(30, True)) - draw.text((96, 302), f"snippet around the answer · {dump['chunk_chars']:,} chars → {dump['chunk_text_tokens']:,} text tokens for the whole chunk", fill=PALETTE["muted"], font=ui_font(18)) + draw.rounded_rectangle( + lane, radius=28, fill=PALETTE["panel"], outline=(35, 49, 59), width=1 + ) + draw.text( + (96, 262), + "text carrier — BPE tokens", + fill=PALETTE["cyan"], + font=ui_font(30, True), + ) + draw.text( + (96, 302), + f"snippet around the answer · {dump['chunk_chars']:,} chars → {dump['chunk_text_tokens']:,} text tokens for the whole chunk", + fill=PALETTE["muted"], + font=ui_font(18), + ) # Token ribbon: show tokens around the answer. tokens = dump["tokens"] @@ -123,27 +171,60 @@ def main() -> None: y += 96 color = PALETTE["amber"] if t["answer"] else (30, 41, 50) text_color = (8, 10, 12) if t["answer"] else PALETTE["ink"] - draw.rounded_rectangle((x, y, x + tw, y + 44), radius=9, fill=color, outline=(52, 68, 80), width=1) + draw.rounded_rectangle( + (x, y, x + tw, y + 44), radius=9, fill=color, outline=(52, 68, 80), width=1 + ) draw.text((x + 11, y + 9), label, fill=text_color, font=fnt) draw.text((x + 4, y + 50), f"id {t['id']}", fill=PALETTE["muted"], font=fnt_id) x += tw + 8 - draw.text((96, 500), "what actually enters the decoder (embedding row, first 6 of " - f"{dump['embed_dim']} dims):", fill=PALETTE["muted"], font=ui_font(18, True)) + draw.text( + (96, 500), + "what actually enters the decoder (embedding row, first 6 of " + f"{dump['embed_dim']} dims):", + fill=PALETTE["muted"], + font=ui_font(18, True), + ) ex = 96 for entry in dump["text_entry"][:3]: box = (ex, 536, ex + 470, 668) - draw.rounded_rectangle(box, radius=16, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1) - draw.text((ex + 18, 548), f"“{entry['str']}” id {entry['id']}", fill=PALETTE["cyan"], font=ui_font(20, True)) - draw.text((ex + 18, 578), vector_text(entry["vector_head"]), fill=PALETTE["ink"], font=mono_font(15)) - draw_vector_bar(draw, (ex + 18, 606), entry["vector_head"], PALETTE["cyan"], width=430) - draw.text((ex + 360, 548), f"‖x‖={entry['norm']:.2f}", fill=PALETTE["muted"], font=ui_font(14)) + draw.rounded_rectangle( + box, radius=16, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1 + ) + draw.text( + (ex + 18, 548), + f"“{entry['str']}” id {entry['id']}", + fill=PALETTE["cyan"], + font=ui_font(20, True), + ) + draw.text( + (ex + 18, 578), + vector_text(entry["vector_head"]), + fill=PALETTE["ink"], + font=mono_font(15), + ) + draw_vector_bar( + draw, (ex + 18, 606), entry["vector_head"], PALETTE["cyan"], width=430 + ) + draw.text( + (ex + 360, 548), + f"‖x‖={entry['norm']:.2f}", + fill=PALETTE["muted"], + font=ui_font(14), + ) ex += 494 # ---- IMAGE LANE ---- lane = (64, 736, 2136, 1336) - draw.rounded_rectangle(lane, radius=28, fill=PALETTE["panel"], outline=(35, 49, 59), width=1) - draw.text((96, 760), "image carrier — visual patch tokens", fill=PALETTE["orange"], font=ui_font(30, True)) + draw.rounded_rectangle( + lane, radius=28, fill=PALETTE["panel"], outline=(35, 49, 59), width=1 + ) + draw.text( + (96, 760), + "image carrier — visual patch tokens", + fill=PALETTE["orange"], + font=ui_font(30, True), + ) px = dump["token_pixel_size"] draw.text( (96, 800), @@ -166,9 +247,16 @@ def main() -> None: cy1 = min(rh, (max(rows) + 1 + pad) * px) crop = resized.crop((cx0, cy0, cx1, cy1)) scale = min(940 / crop.width, 225 / crop.height) - crop_big = crop.resize((round(crop.width * scale), round(crop.height * scale)), Image.Resampling.NEAREST) + crop_big = crop.resize( + (round(crop.width * scale), round(crop.height * scale)), + Image.Resampling.NEAREST, + ) ox, oy = 96, 852 - draw.rounded_rectangle((ox - 6, oy - 6, ox + crop_big.width + 6, oy + crop_big.height + 6), radius=10, fill=(244, 242, 230)) + draw.rounded_rectangle( + (ox - 6, oy - 6, ox + crop_big.width + 6, oy + crop_big.height + 6), + radius=10, + fill=(244, 242, 230), + ) canvas.paste(crop_big, (ox, oy)) cd = ImageDraw.Draw(canvas) for gx in range(cx0 // px, cx1 // px + 1): @@ -181,39 +269,111 @@ def main() -> None: r, c = idx // grid, idx % grid xa = ox + (c * px - cx0) * scale ya = oy + (r * px - cy0) * scale - cd.rectangle((xa, ya, xa + px * scale, ya + px * scale), outline=PALETTE["orange"], width=4) - draw.text((ox, oy + crop_big.height + 14), f"orange cells = the {len(indices)} visual tokens covering “{answer}” (token grid {grid}×{grid})", fill=PALETTE["muted"], font=ui_font(17)) + cd.rectangle( + (xa, ya, xa + px * scale, ya + px * scale), + outline=PALETTE["orange"], + width=4, + ) + draw.text( + (ox, oy + crop_big.height + 14), + f"orange cells = the {len(indices)} visual tokens covering “{answer}” (token grid {grid}×{grid})", + fill=PALETTE["muted"], + font=ui_font(17), + ) # Magnified single patches. sx = ox + crop_big.width + 60 - draw.text((sx, 852 - 26), "individual visual tokens (real input pixels):", fill=PALETTE["muted"], font=ui_font(18, True)) + draw.text( + (sx, 852 - 26), + "individual visual tokens (real input pixels):", + fill=PALETTE["muted"], + font=ui_font(18, True), + ) for k, idx in enumerate(indices[:5]): r, c = idx // grid, idx % grid - cell = resized.crop((c * px, r * px, (c + 1) * px, (r + 1) * px)).resize((132, 132), Image.Resampling.NEAREST) + cell = resized.crop((c * px, r * px, (c + 1) * px, (r + 1) * px)).resize( + (132, 132), Image.Resampling.NEAREST + ) bx = sx + k * 160 - draw.rounded_rectangle((bx - 4, 852 - 4, bx + 136, 852 + 136), radius=8, fill=(244, 242, 230), outline=PALETTE["orange"], width=3) + draw.rounded_rectangle( + (bx - 4, 852 - 4, bx + 136, 852 + 136), + radius=8, + fill=(244, 242, 230), + outline=PALETTE["orange"], + width=3, + ) canvas.paste(cell, (bx, 852)) draw.text((bx, 996), f"tok[{idx}]", fill=PALETTE["muted"], font=mono_font(13)) - draw.text((sx, 1030), f"pre-tower normalized pixels of first patch: {vector_text(dump['pixel_head_first_answer_patch'])}", fill=PALETTE["muted"], font=mono_font(14)) + draw.text( + (sx, 1030), + f"pre-tower normalized pixels of first patch: {vector_text(dump['pixel_head_first_answer_patch'])}", + fill=PALETTE["muted"], + font=mono_font(14), + ) - draw.text((96, 1106), f"what actually enters the decoder (visual-tower output, first 6 of {dump['visual_out_dim']} dims):", fill=PALETTE["muted"], font=ui_font(18, True)) + draw.text( + (96, 1106), + f"what actually enters the decoder (visual-tower output, first 6 of {dump['visual_out_dim']} dims):", + fill=PALETTE["muted"], + font=ui_font(18, True), + ) ex = 96 for entry in dump["image_entry"][:4]: box = (ex, 1142, ex + 470, 1274) - draw.rounded_rectangle(box, radius=16, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1) + draw.rounded_rectangle( + box, radius=16, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1 + ) r, c = entry["grid_rc"] - draw.text((ex + 18, 1154), f"visual tok[{entry['token_index']}] (row {r}, col {c})", fill=PALETTE["orange"], font=ui_font(20, True)) - draw.text((ex + 18, 1184), vector_text(entry["vector_head"]), fill=PALETTE["ink"], font=mono_font(15)) - draw_vector_bar(draw, (ex + 18, 1212), entry["vector_head"], PALETTE["orange"], width=430) - draw.text((ex + 360, 1154), f"‖x‖={entry['norm']:.2f}", fill=PALETTE["muted"], font=ui_font(14)) + draw.text( + (ex + 18, 1154), + f"visual tok[{entry['token_index']}] (row {r}, col {c})", + fill=PALETTE["orange"], + font=ui_font(20, True), + ) + draw.text( + (ex + 18, 1184), + vector_text(entry["vector_head"]), + fill=PALETTE["ink"], + font=mono_font(15), + ) + draw_vector_bar( + draw, (ex + 18, 1212), entry["vector_head"], PALETTE["orange"], width=430 + ) + draw.text( + (ex + 360, 1154), + f"‖x‖={entry['norm']:.2f}", + fill=PALETTE["muted"], + font=ui_font(14), + ) ex += 494 # Comparison strip. text_tok_for_word = len(dump["text_entry"]) - draw.rounded_rectangle((1100, 536, 2104, 668), radius=16, fill=PALETTE["panel2"], outline=(34, 48, 58), width=1) - draw.text((1128, 556), f"“{answer}” = {text_tok_for_word} text token(s) · {len(indices)} visual tokens", fill=PALETTE["ink"], font=ui_font(22, True)) - draw.text((1128, 592), f"both end up as {dump['embed_dim']}-dim rows in the same decoder", fill=PALETTE["ink"], font=ui_font(19)) - draw.text((1128, 626), "text path: lookup table row. image path: ViT forward over 4 raw patches → merger MLP.", fill=PALETTE["muted"], font=ui_font(16)) + draw.rounded_rectangle( + (1100, 536, 2104, 668), + radius=16, + fill=PALETTE["panel2"], + outline=(34, 48, 58), + width=1, + ) + draw.text( + (1128, 556), + f"“{answer}” = {text_tok_for_word} text token(s) · {len(indices)} visual tokens", + fill=PALETTE["ink"], + font=ui_font(22, True), + ) + draw.text( + (1128, 592), + f"both end up as {dump['embed_dim']}-dim rows in the same decoder", + fill=PALETTE["ink"], + font=ui_font(19), + ) + draw.text( + (1128, 626), + "text path: lookup table row. image path: ViT forward over 4 raw patches → merger MLP.", + fill=PALETTE["muted"], + font=ui_font(16), + ) out = Path(args.out) out.parent.mkdir(parents=True, exist_ok=True) diff --git a/packages/snapcompact/research/snapcompact_viz_atlas.py b/packages/snapcompact/research/snapcompact_viz_atlas.py index 29d05c57b..b28794780 100644 --- a/packages/snapcompact/research/snapcompact_viz_atlas.py +++ b/packages/snapcompact/research/snapcompact_viz_atlas.py @@ -39,9 +39,13 @@ PURPLE = (183, 108, 255) def font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: candidates = [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", "/System/Library/Fonts/Supplemental/Avenir Next Condensed.ttc", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ] for path in candidates: if path and Path(path).exists(): @@ -82,11 +86,16 @@ def normalize_coords(coords: np.ndarray) -> np.ndarray: return out -def kmeans(points: np.ndarray, k: int = 5, iters: int = 32) -> tuple[np.ndarray, np.ndarray]: +def kmeans( + points: np.ndarray, k: int = 5, iters: int = 32 +) -> tuple[np.ndarray, np.ndarray]: # Deterministic farthest-point seeding avoids random output drift. centers = [points[np.argmax(points[:, 0] + points[:, 1])]] for _ in range(1, k): - dist = np.min(np.sum((points[:, None, :] - np.asarray(centers)[None, :, :]) ** 2, axis=2), axis=1) + dist = np.min( + np.sum((points[:, None, :] - np.asarray(centers)[None, :, :]) ** 2, axis=2), + axis=1, + ) centers.append(points[int(np.argmax(dist))]) c = np.asarray(centers, dtype=np.float32) labels = np.zeros(points.shape[0], dtype=np.int32) @@ -103,7 +112,9 @@ def kmeans(points: np.ndarray, k: int = 5, iters: int = 32) -> tuple[np.ndarray, return labels, c -def crop_answer_strip(img: Image.Image, start: int, end: int, cols: int, adv: int = 8, pitch: int = 13) -> Image.Image: +def crop_answer_strip( + img: Image.Image, start: int, end: int, cols: int, adv: int = 8, pitch: int = 13 +) -> Image.Image: row0 = max(0, start // cols - 4) row1 = min(img.height // pitch, end // cols + 5) col0 = max(0, start % cols - 32) @@ -118,11 +129,19 @@ def crop_answer_strip(img: Image.Image, start: int, end: int, cols: int, adv: in return crop -def paste_fit(canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int]) -> None: +def paste_fit( + canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int] +) -> None: x0, y0, x1, y1 = box scale = min((x1 - x0) / img.width, (y1 - y0) / img.height) - resized = img.resize((max(1, round(img.width * scale)), max(1, round(img.height * scale))), Image.Resampling.NEAREST) - canvas.paste(resized, (x0 + (x1 - x0 - resized.width) // 2, y0 + (y1 - y0 - resized.height) // 2)) + resized = img.resize( + (max(1, round(img.width * scale)), max(1, round(img.height * scale))), + Image.Resampling.NEAREST, + ) + canvas.paste( + resized, + (x0 + (x1 - x0 - resized.width) // 2, y0 + (y1 - y0 - resized.height) // 2), + ) def render_atlas_panel( @@ -136,14 +155,26 @@ def render_atlas_panel( summary: dict, out_dir: Path, ) -> Image.Image: - cmap = LinearSegmentedColormap.from_list("scar", ["#182132", "#245d7a", "#48d8ff", "#ffd04e", "#ff493d"]) + cmap = LinearSegmentedColormap.from_list( + "scar", ["#182132", "#245d7a", "#48d8ff", "#ffd04e", "#ff493d"] + ) fig = plt.figure(figsize=(15.8, 10.6), dpi=170) fig.patch.set_facecolor("#04070c") ax = fig.add_axes((0.045, 0.06, 0.91, 0.88), facecolor="#07101a") x = points[:, 0] y = points[:, 1] - hb = ax.hexbin(x, y, C=ratio_strength, gridsize=46, reduce_C_function=np.mean, cmap=cmap, mincnt=1, linewidths=0, alpha=0.64) + hb = ax.hexbin( + x, + y, + C=ratio_strength, + gridsize=46, + reduce_C_function=np.mean, + cmap=cmap, + mincnt=1, + linewidths=0, + alpha=0.64, + ) hb.set_clim(0.0, 1.0) cluster_colors = ["#46d8ff", "#ff4b3d", "#ffc644", "#87ff8b", "#b76cff"] @@ -151,7 +182,14 @@ def render_atlas_panel( mask = labels == i if np.count_nonzero(mask) < 4: continue - ax.scatter(x[mask], y[mask], s=28 + answer_strength[mask] * 150, c=color, alpha=0.24, linewidths=0) + ax.scatter( + x[mask], + y[mask], + s=28 + answer_strength[mask] * 150, + c=color, + alpha=0.24, + linewidths=0, + ) ax.scatter( x[mask], y[mask], @@ -165,7 +203,15 @@ def render_atlas_panel( ) hot = np.argsort(ratio_strength + answer_strength * 0.55)[-9:] - ax.scatter(x[hot], y[hot], s=210, facecolors="none", edgecolors="#fff0a8", linewidths=1.5, alpha=0.95) + ax.scatter( + x[hot], + y[hot], + s=210, + facecolors="none", + edgecolors="#fff0a8", + linewidths=1.5, + alpha=0.95, + ) for rank, idx in enumerate(hot[-5:][::-1], 1): ax.text( x[idx] + 0.012, @@ -177,12 +223,23 @@ def render_atlas_panel( path_effects=[pe.withStroke(linewidth=2.5, foreground="#05070a")], ) - names = ["answer ridge", "control basin", "early glyph shore", "late-context upland", "ratio reef"] + names = [ + "answer ridge", + "control basin", + "early glyph shore", + "late-context upland", + "ratio reef", + ] cluster_scores = [] for i in range(len(centers)): mask = labels == i - cluster_scores.append((float(ratio_strength[mask].mean()) if np.any(mask) else 0.0, i)) - order = {old: new for new, (_score, old) in enumerate(sorted(cluster_scores, reverse=True))} + cluster_scores.append( + (float(ratio_strength[mask].mean()) if np.any(mask) else 0.0, i) + ) + order = { + old: new + for new, (_score, old) in enumerate(sorted(cluster_scores, reverse=True)) + } for i, c in enumerate(centers): mask = labels == i if np.count_nonzero(mask) < 5: @@ -277,14 +334,28 @@ def draw_shell(panel: Image.Image, summary: dict, data_dir: Path, out: Path) -> gd.ellipse((-360, -240, 1000, 760), fill=(70, 216, 255, 32)) gd.ellipse((1270, 210, 2740, 1610), fill=(255, 75, 61, 32)) gd.ellipse((690, 920, 1740, 1780), fill=(255, 198, 68, 18)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(90))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(90)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) draw.text((74, 48), "SNAPCOMPACT WHITEBOX", fill=AMBER, font=font(24, True)) - draw.text((74, 88), "Activation Atlas of the missing answer", fill=INK, font=font(72, True)) - draw.text((78, 178), "A PCA geography of image-token residual scars: where blanking the gold answer ‘2003’ moves the model differently than a random blank.", fill=MUTED, font=font(27)) + draw.text( + (74, 88), + "Activation Atlas of the missing answer", + fill=INK, + font=font(72, True), + ) + draw.text( + (78, 178), + "A PCA geography of image-token residual scars: where blanking the gold answer ‘2003’ moves the model differently than a random blank.", + fill=MUTED, + font=font(27), + ) - draw.rounded_rectangle((74, 252, 590, 1390), radius=32, fill=PANEL, outline=(32, 45, 58), width=1) + draw.rounded_rectangle( + (74, 252, 590, 1390), radius=32, fill=PANEL, outline=(32, 45, 58), width=1 + ) q = summary["question"] cols = summary["geometry"]["cols"] original = Image.open(data_dir / "images" / "original.png").convert("RGB") @@ -298,11 +369,27 @@ def draw_shell(panel: Image.Image, summary: dict, data_dir: Path, out: Path) -> y = 314 for title, img, color in strips: draw.text((110, y), title, fill=color, font=font(18, True)) - draw.rounded_rectangle((110, y + 28, 554, y + 166), radius=16, fill=(242, 241, 229), outline=color, width=3) - paste_fit(canvas, crop_answer_strip(img, q["answer_start"], q["answer_end"], cols), (124, y + 42, 540, y + 152)) + draw.rounded_rectangle( + (110, y + 28, 554, y + 166), + radius=16, + fill=(242, 241, 229), + outline=color, + width=3, + ) + paste_fit( + canvas, + crop_answer_strip(img, q["answer_start"], q["answer_end"], cols), + (124, y + 42, 540, y + 152), + ) y += 226 - draw.rounded_rectangle((110, 1002, 554, 1300), radius=24, fill=(7, 11, 18), outline=(35, 51, 66), width=1) + draw.rounded_rectangle( + (110, 1002, 554, 1300), + radius=24, + fill=(7, 11, 18), + outline=(35, 51, 66), + width=1, + ) metrics = [ ("answer", q["answer_text"], AMBER, 48), ("layers", str(summary["layers"]), CYAN, 38), @@ -314,9 +401,16 @@ def draw_shell(panel: Image.Image, summary: dict, data_dir: Path, out: Path) -> draw.text((142, yy), label, fill=MUTED, font=font(16, True)) draw.text((142, yy + 26), value, fill=color, font=font(size, True)) yy += 68 - draw.text((110, 1336), "Actual heatmaps.npz + summary.json; no schematic points.", fill=MUTED, font=font(18)) + draw.text( + (110, 1336), + "Actual heatmaps.npz + summary.json; no schematic points.", + fill=MUTED, + font=font(18), + ) - draw.rounded_rectangle((622, 252, 2326, 1390), radius=32, fill=PANEL, outline=(32, 45, 58), width=1) + draw.rounded_rectangle( + (622, 252, 2326, 1390), radius=32, fill=PANEL, outline=(32, 45, 58), width=1 + ) panel = panel.resize((1640, 1098), Image.Resampling.LANCZOS) canvas.paste(panel, (654, 272)) @@ -324,7 +418,15 @@ def draw_shell(panel: Image.Image, summary: dict, data_dir: Path, out: Path) -> canvas.save(out, quality=95) -def write_source_data(out_dir: Path, points: np.ndarray, labels: np.ndarray, ratio_strength: np.ndarray, answer_strength: np.ndarray, peak_layers: np.ndarray, explained: np.ndarray) -> None: +def write_source_data( + out_dir: Path, + points: np.ndarray, + labels: np.ndarray, + ratio_strength: np.ndarray, + answer_strength: np.ndarray, + peak_layers: np.ndarray, + explained: np.ndarray, +) -> None: np.savez_compressed( out_dir / "atlas_source.npz", points=points, @@ -336,9 +438,29 @@ def write_source_data(out_dir: Path, points: np.ndarray, labels: np.ndarray, rat ) with (out_dir / "atlas_points.csv").open("w", newline="") as f: writer = csv.writer(f) - writer.writerow(["token", "atlas_x", "atlas_y", "cluster", "ratio_strength", "answer_strength", "peak_layer"]) + writer.writerow( + [ + "token", + "atlas_x", + "atlas_y", + "cluster", + "ratio_strength", + "answer_strength", + "peak_layer", + ] + ) for i in range(points.shape[0]): - writer.writerow([i, f"{points[i, 0]:.6f}", f"{points[i, 1]:.6f}", int(labels[i]), f"{ratio_strength[i]:.6f}", f"{answer_strength[i]:.6f}", int(peak_layers[i])]) + writer.writerow( + [ + i, + f"{points[i, 0]:.6f}", + f"{points[i, 1]:.6f}", + int(labels[i]), + f"{ratio_strength[i]:.6f}", + f"{answer_strength[i]:.6f}", + int(peak_layers[i]), + ] + ) def main() -> None: @@ -357,7 +479,9 @@ def main() -> None: ratio = heatmaps["ratio"].astype(np.float32) contrast = np.log1p(answer) - np.log1p(random) - features = np.concatenate([contrast.T, np.log1p(ratio).T, np.log1p(answer).T], axis=1) + features = np.concatenate( + [contrast.T, np.log1p(ratio).T, np.log1p(answer).T], axis=1 + ) raw_coords, explained = pca2(features) points = normalize_coords(raw_coords) labels, centers = kmeans(points, k=5) @@ -366,8 +490,20 @@ def main() -> None: answer_strength = quantile_norm(np.log1p(answer).mean(axis=0), 0.02, 0.99) peak_layers = np.argmax(ratio, axis=0).astype(np.int32) - write_source_data(out_dir, points, labels, ratio_strength, answer_strength, peak_layers, explained) - panel = render_atlas_panel(points, labels, centers, ratio_strength, answer_strength, peak_layers, explained, summary, out_dir) + write_source_data( + out_dir, points, labels, ratio_strength, answer_strength, peak_layers, explained + ) + panel = render_atlas_panel( + points, + labels, + centers, + ratio_strength, + answer_strength, + peak_layers, + explained, + summary, + out_dir, + ) out = out_dir / "atlas.png" draw_shell(panel, summary, data_dir, out) print(out) diff --git a/packages/snapcompact/research/snapcompact_viz_circuit.py b/packages/snapcompact/research/snapcompact_viz_circuit.py index f78c70df3..57fca20e8 100644 --- a/packages/snapcompact/research/snapcompact_viz_circuit.py +++ b/packages/snapcompact/research/snapcompact_viz_circuit.py @@ -33,9 +33,15 @@ GREEN = (127, 245, 148) def font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: candidates = [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/System/Library/Fonts/Supplemental/Helvetica Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Helvetica.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/System/Library/Fonts/Supplemental/Helvetica Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Helvetica.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ] for path in candidates: if Path(path).exists(): @@ -47,7 +53,9 @@ def clamp01(v: float) -> float: return max(0.0, min(1.0, v)) -def mix(a: tuple[int, int, int], b: tuple[int, int, int], t: float) -> tuple[int, int, int]: +def mix( + a: tuple[int, int, int], b: tuple[int, int, int], t: float +) -> tuple[int, int, int]: t = clamp01(t) return tuple(round(x + (y - x) * t) for x, y in zip(a, b)) @@ -59,11 +67,24 @@ def quantile_norm(values: np.ndarray, q: float = 0.97) -> np.ndarray: return np.clip(values / scale, 0, 1) -def rounded_panel(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], radius: int = 34) -> None: - draw.rounded_rectangle(box, radius=radius, fill=PANEL, outline=(31, 41, 51), width=1) +def rounded_panel( + draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], radius: int = 34 +) -> None: + draw.rounded_rectangle( + box, radius=radius, fill=PANEL, outline=(31, 41, 51), width=1 + ) -def multiline(draw: ImageDraw.ImageDraw, xy: tuple[int, int], text: str, *, fill: tuple[int, int, int], fnt: ImageFont.ImageFont, max_width: int, line_gap: int = 8) -> int: +def multiline( + draw: ImageDraw.ImageDraw, + xy: tuple[int, int], + text: str, + *, + fill: tuple[int, int, int], + fnt: ImageFont.ImageFont, + max_width: int, + line_gap: int = 8, +) -> int: words = text.split() lines: list[str] = [] cur = "" @@ -85,7 +106,9 @@ def multiline(draw: ImageDraw.ImageDraw, xy: tuple[int, int], text: str, *, fill return y -def crop_answer_region(img: Image.Image, summary: dict, pad_cells: int = 42) -> Image.Image: +def crop_answer_region( + img: Image.Image, summary: dict, pad_cells: int = 42 +) -> Image.Image: q = summary["question"] cols = int(summary["geometry"]["cols"]) rows = int(summary["geometry"]["rows"]) @@ -109,15 +132,34 @@ def crop_answer_region(img: Image.Image, summary: dict, pad_cells: int = 42) -> return crop -def paste_fit(canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int], *, resample: int = Image.Resampling.LANCZOS) -> None: +def paste_fit( + canvas: Image.Image, + img: Image.Image, + box: tuple[int, int, int, int], + *, + resample: int = Image.Resampling.LANCZOS, +) -> None: x0, y0, x1, y1 = box scale = min((x1 - x0) / img.width, (y1 - y0) / img.height) size = (max(1, round(img.width * scale)), max(1, round(img.height * scale))) resized = img.resize(size, resample) - canvas.paste(resized, (x0 + (x1 - x0 - size[0]) // 2, y0 + (y1 - y0 - size[1]) // 2)) + canvas.paste( + resized, (x0 + (x1 - x0 - size[0]) // 2, y0 + (y1 - y0 - size[1]) // 2) + ) -def draw_bezier(draw: ImageDraw.ImageDraw, points: tuple[tuple[float, float], tuple[float, float], tuple[float, float], tuple[float, float]], *, fill: tuple[int, int, int, int], width: int) -> None: +def draw_bezier( + draw: ImageDraw.ImageDraw, + points: tuple[ + tuple[float, float], + tuple[float, float], + tuple[float, float], + tuple[float, float], + ], + *, + fill: tuple[int, int, int, int], + width: int, +) -> None: p0, p1, p2, p3 = points coords: list[tuple[float, float]] = [] for i in range(46): @@ -131,7 +173,17 @@ def draw_bezier(draw: ImageDraw.ImageDraw, points: tuple[tuple[float, float], tu def token_groups(grid_side: int = 27, tiles: int = 3) -> list[dict[str, int | str]]: groups: list[dict[str, int | str]] = [] - names = ["upper-left", "upper", "upper-right", "left", "center", "right", "lower-left", "lower", "lower-right"] + names = [ + "upper-left", + "upper", + "upper-right", + "left", + "center", + "right", + "lower-left", + "lower", + "lower-right", + ] idx = 0 for gy in range(tiles): y0 = round(gy * grid_side / tiles) @@ -152,7 +204,9 @@ def group_indices(group: dict[str, int | str], grid_side: int = 27) -> np.ndarra return np.asarray(ids, dtype=np.int64) -def build_metrics(answer: np.ndarray, random: np.ndarray, ratio: np.ndarray) -> tuple[list[dict], list[dict], np.ndarray, np.ndarray]: +def build_metrics( + answer: np.ndarray, random: np.ndarray, ratio: np.ndarray +) -> tuple[list[dict], list[dict], np.ndarray, np.ndarray]: layers, tokens = answer.shape grid_side = int(round(math.sqrt(tokens))) if grid_side * grid_side != tokens: @@ -195,14 +249,21 @@ def build_metrics(answer: np.ndarray, random: np.ndarray, ratio: np.ndarray) -> "answer_delta_mean": float(answer[layer].mean()), "random_delta_mean": float(random[layer].mean()), "ratio_mean": float(ratio[layer].mean()), - "answer_minus_random_mean": float((answer[layer] - random[layer]).mean()), + "answer_minus_random_mean": float( + (answer[layer] - random[layer]).mean() + ), "edge_score": float(layer_score[layer].mean()), } ) return group_rows, layer_rows, layer_score, layer_norm -def draw_token_grid(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], token_strength: np.ndarray, group_rows: list[dict]) -> list[tuple[int, int]]: +def draw_token_grid( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + token_strength: np.ndarray, + group_rows: list[dict], +) -> list[tuple[int, int]]: x0, y0, x1, y1 = box grid = token_strength.reshape(27, 27) norm = quantile_norm(grid, 0.985) @@ -215,8 +276,18 @@ def draw_token_grid(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], t color = mix((13, 24, 34), ORANGE, v) if v > 0.72: color = mix(color, GOLD, (v - 0.72) / 0.28) - draw.rectangle((gx + x * cell, gy + y * cell, gx + (x + 1) * cell - 1, gy + (y + 1) * cell - 1), fill=color) - draw.rectangle((gx - 1, gy - 1, gx + 27 * cell, gy + 27 * cell), outline=(70, 88, 101), width=2) + draw.rectangle( + ( + gx + x * cell, + gy + y * cell, + gx + (x + 1) * cell - 1, + gy + (y + 1) * cell - 1, + ), + fill=color, + ) + draw.rectangle( + (gx - 1, gy - 1, gx + 27 * cell, gy + 27 * cell), outline=(70, 88, 101), width=2 + ) centers: list[tuple[int, int]] = [] scores = np.asarray([g["edge_score"] for g in group_rows], dtype=np.float32) @@ -226,11 +297,17 @@ def draw_token_grid(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], t cy = gy + round((int(g["y0"]) + int(g["y1"])) * 0.5 * cell) centers.append((cx, cy)) rad = round(8 + 19 * float(s)) - draw.ellipse((cx - rad, cy - rad, cx + rad, cy + rad), outline=mix(BLUE, GOLD, float(s)), width=3) + draw.ellipse( + (cx - rad, cy - rad, cx + rad, cy + rad), + outline=mix(BLUE, GOLD, float(s)), + width=3, + ) return centers -def draw_layer_bands(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], layer_rows: list[dict]) -> list[tuple[int, int]]: +def draw_layer_bands( + draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], layer_rows: list[dict] +) -> list[tuple[int, int]]: x0, y0, x1, y1 = box scores = np.asarray([r["edge_score"] for r in layer_rows], dtype=np.float32) ratios = np.asarray([r["ratio_mean"] for r in layer_rows], dtype=np.float32) @@ -245,14 +322,34 @@ def draw_layer_bands(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], inset = round(26 * (1 - float(s))) color = mix((17, 27, 38), GOLD, float(rr) * 0.80) outline = mix((54, 71, 83), RED, float(s)) - draw.rounded_rectangle((x0 + inset, yy0, x1 - inset, yy1), radius=8, fill=color, outline=outline, width=2) - draw.text((x0 - 74, yy0 + max(0, (yy1 - yy0 - 18) // 2)), f"L{int(row['layer']):02d}", fill=mix(MUTED, INK, float(s)), font=font(16, True)) + draw.rounded_rectangle( + (x0 + inset, yy0, x1 - inset, yy1), + radius=8, + fill=color, + outline=outline, + width=2, + ) + draw.text( + (x0 - 74, yy0 + max(0, (yy1 - yy0 - 18) // 2)), + f"L{int(row['layer']):02d}", + fill=mix(MUTED, INK, float(s)), + font=font(16, True), + ) centers.append(((x0 + x1) // 2, (yy0 + yy1) // 2)) return centers -def render(summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np.ndarray, result_dir: Path, out_dir: Path) -> dict: - group_rows, layer_rows, layer_score, layer_norm = build_metrics(answer, random, ratio) +def render( + summary: dict, + answer: np.ndarray, + random: np.ndarray, + ratio: np.ndarray, + result_dir: Path, + out_dir: Path, +) -> dict: + group_rows, layer_rows, layer_score, layer_norm = build_metrics( + answer, random, ratio + ) w, h = 2400, 1350 canvas = Image.new("RGB", (w, h), BG) draw = ImageDraw.Draw(canvas) @@ -264,11 +361,18 @@ def render(summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np.ndar gd.ellipse((-380, -260, 980, 640), fill=(255, 72, 82, 36)) gd.ellipse((780, 70, 2320, 1420), fill=(83, 218, 255, 22)) gd.ellipse((1440, -120, 2760, 860), fill=(255, 199, 74, 26)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(90))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(90)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) draw.text((70, 44), "SNAPCOMPACT CIRCUIT TRACE", fill=GOLD, font=font(25, True)) - draw.text((70, 82), "The answer glyphs light a decoder circuit", fill=INK, font=font(66, True)) + draw.text( + (70, 82), + "The answer glyphs light a decoder circuit", + fill=INK, + font=font(66, True), + ) multiline( draw, (72, 166), @@ -292,59 +396,133 @@ def render(summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np.ndar draw.text((94, 294), "1. bitmap intervention", fill=INK, font=font(30, True)) draw.text((94, 334), "question targets one visible year", fill=MUTED, font=font(18)) - draw.rounded_rectangle((94, 386, 488, 560), radius=16, fill=(240, 238, 224), outline=BLUE, width=3) + draw.rounded_rectangle( + (94, 386, 488, 560), radius=16, fill=(240, 238, 224), outline=BLUE, width=3 + ) paste_fit(canvas, crop, (108, 400, 474, 546), resample=Image.Resampling.NEAREST) draw.text((94, 574), "original answer region", fill=BLUE, font=font(18, True)) - draw.rounded_rectangle((94, 654, 488, 828), radius=16, fill=(240, 238, 224), outline=RED, width=3) - paste_fit(canvas, masked_crop, (108, 668, 474, 814), resample=Image.Resampling.NEAREST) + draw.rounded_rectangle( + (94, 654, 488, 828), radius=16, fill=(240, 238, 224), outline=RED, width=3 + ) + paste_fit( + canvas, masked_crop, (108, 668, 474, 814), resample=Image.Resampling.NEAREST + ) draw.text((94, 842), "blanked answer mask", fill=RED, font=font(18, True)) draw.text((94, 930), "question", fill=MUTED, font=font(15, True)) - multiline(draw, (94, 956), str(q["q"]), fill=INK, fnt=font(23), max_width=370, line_gap=8) + multiline( + draw, (94, 956), str(q["q"]), fill=INK, fnt=font(23), max_width=370, line_gap=8 + ) draw.text((94, 1070), "gold answer", fill=MUTED, font=font(15, True)) draw.text((94, 1098), str(q["answer_text"]), fill=GOLD, font=font(52, True)) - draw.text((94, 1172), f"global Δ ratio {summary['answer_over_random_delta']:.2f}×", fill=INK, font=font(22, True)) + draw.text( + (94, 1172), + f"global Δ ratio {summary['answer_over_random_delta']:.2f}×", + fill=INK, + font=font(22, True), + ) draw.text((600, 294), "2. image-token regions", fill=INK, font=font(30, True)) - draw.text((600, 334), "27×27 token lattice, colored by circuit score", fill=MUTED, font=font(18)) + draw.text( + (600, 334), + "27×27 token lattice, colored by circuit score", + fill=MUTED, + font=font(18), + ) token_strength = layer_score.mean(axis=0) - token_centers = draw_token_grid(draw, (616, 392, 954, 730), token_strength, group_rows) + token_centers = draw_token_grid( + draw, (616, 392, 954, 730), token_strength, group_rows + ) top_groups = sorted(group_rows, key=lambda g: g["edge_score"], reverse=True)[:4] draw.text((600, 794), "strongest token groups", fill=MUTED, font=font(16, True)) y = 826 - group_score_norm = quantile_norm(np.asarray([g["edge_score"] for g in group_rows], dtype=np.float32), 0.92) + group_score_norm = quantile_norm( + np.asarray([g["edge_score"] for g in group_rows], dtype=np.float32), 0.92 + ) for g in top_groups: s = float(group_score_norm[int(g["id"])]) - draw.rounded_rectangle((600, y, 970, y + 62), radius=14, fill=PANEL_2, outline=mix((44, 58, 68), GOLD, s), width=2) + draw.rounded_rectangle( + (600, y, 970, y + 62), + radius=14, + fill=PANEL_2, + outline=mix((44, 58, 68), GOLD, s), + width=2, + ) draw.text((620, y + 12), str(g["name"]), fill=INK, font=font(20, True)) - draw.text((820, y + 12), f"{g['ratio_mean']:.2f}×", fill=mix(BLUE, GOLD, s), font=font(21, True)) - draw.text((620, y + 38), f"Δ {g['answer_delta_mean']:.2f} vs {g['random_delta_mean']:.2f}", fill=MUTED, font=font(14)) + draw.text( + (820, y + 12), + f"{g['ratio_mean']:.2f}×", + fill=mix(BLUE, GOLD, s), + font=font(21, True), + ) + draw.text( + (620, y + 38), + f"Δ {g['answer_delta_mean']:.2f} vs {g['random_delta_mean']:.2f}", + fill=MUTED, + font=font(14), + ) y += 78 draw.text((1158, 294), "3. decoder layer bands", fill=INK, font=font(30, True)) - draw.text((1158, 334), "band width/color follows per-layer answer specificity", fill=MUTED, font=font(18)) + draw.text( + (1158, 334), + "band width/color follows per-layer answer specificity", + fill=MUTED, + font=font(18), + ) layer_centers = draw_layer_bands(draw, (1246, 394, 1566, 1122), layer_rows) draw.text((1848, 294), "4. output answer", fill=INK, font=font(30, True)) - draw.text((1848, 334), "residual stream converges on text", fill=MUTED, font=font(18)) - draw.rounded_rectangle((1880, 462, 2274, 730), radius=34, fill=(10, 13, 18), outline=(73, 82, 92), width=2) + draw.text( + (1848, 334), "residual stream converges on text", fill=MUTED, font=font(18) + ) + draw.rounded_rectangle( + (1880, 462, 2274, 730), + radius=34, + fill=(10, 13, 18), + outline=(73, 82, 92), + width=2, + ) draw.text((1918, 500), "PaddleOCR-VL", fill=MUTED, font=font(20, True)) draw.text((1918, 558), "answers", fill=INK, font=font(32, True)) draw.text((1918, 606), str(q["answer_text"]), fill=GOLD, font=font(82, True)) - draw.rounded_rectangle((1880, 820, 2274, 1034), radius=28, fill=PANEL_2, outline=(47, 62, 73), width=2) - draw.text((1918, 858), f"{summary['layers']} decoder layers", fill=INK, font=font(26, True)) - draw.text((1918, 900), f"{summary['image_tokens']} image tokens", fill=MUTED, font=font(21)) - draw.text((1918, 938), "edge thickness = grouped delta score", fill=MUTED, font=font(21)) - draw.text((1918, 976), "edge color = answer/random ratio", fill=MUTED, font=font(21)) + draw.rounded_rectangle( + (1880, 820, 2274, 1034), radius=28, fill=PANEL_2, outline=(47, 62, 73), width=2 + ) + draw.text( + (1918, 858), + f"{summary['layers']} decoder layers", + fill=INK, + font=font(26, True), + ) + draw.text( + (1918, 900), + f"{summary['image_tokens']} image tokens", + fill=MUTED, + font=font(21), + ) + draw.text( + (1918, 938), "edge thickness = grouped delta score", fill=MUTED, font=font(21) + ) + draw.text( + (1918, 976), "edge color = answer/random ratio", fill=MUTED, font=font(21) + ) # Edges live in a transparent layer so glow can sit behind node labels. edges = Image.new("RGBA", (w, h), (0, 0, 0, 0)) ed = ImageDraw.Draw(edges) - all_group_layer = np.asarray([g["layer_scores"] for g in group_rows], dtype=np.float32) + all_group_layer = np.asarray( + [g["layer_scores"] for g in group_rows], dtype=np.float32 + ) group_layer_norm = quantile_norm(all_group_layer, 0.965) - group_layer_ratio = np.asarray([g["layer_ratios"] for g in group_rows], dtype=np.float32) + group_layer_ratio = np.asarray( + [g["layer_ratios"] for g in group_rows], dtype=np.float32 + ) ratio_norm = quantile_norm(group_layer_ratio, 0.955) - selected_groups = [int(g["id"]) for g in sorted(group_rows, key=lambda g: g["edge_score"], reverse=True)[:7]] + selected_groups = [ + int(g["id"]) + for g in sorted(group_rows, key=lambda g: g["edge_score"], reverse=True)[:7] + ] selected_layers = [0, 1, 2, 3, 4, 5, 7, 9, 12, 15, 18] for gi in selected_groups: sx, sy = token_centers[gi] @@ -356,10 +534,19 @@ def render(summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np.ndar col = mix(BLUE, RED, float(ratio_norm[gi, li])) alpha = round(54 + 156 * strength) width = max(1, round(1 + 9 * strength)) - draw_bezier(ed, ((sx + 18, sy), (1046, sy), (1110, ey), (ex - 162, ey)), fill=(*col, alpha), width=width) + draw_bezier( + ed, + ((sx + 18, sy), (1046, sy), (1110, ey), (ex - 162, ey)), + fill=(*col, alpha), + width=width, + ) - layer_edge_norm = quantile_norm(np.asarray([r["edge_score"] for r in layer_rows], dtype=np.float32), 0.96) - layer_ratio_norm = quantile_norm(np.asarray([r["ratio_mean"] for r in layer_rows], dtype=np.float32), 0.96) + layer_edge_norm = quantile_norm( + np.asarray([r["edge_score"] for r in layer_rows], dtype=np.float32), 0.96 + ) + layer_ratio_norm = quantile_norm( + np.asarray([r["ratio_mean"] for r in layer_rows], dtype=np.float32), 0.96 + ) out_anchor = (1880, 596) for li in selected_layers: sx, sy = layer_centers[li] @@ -367,7 +554,17 @@ def render(summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np.ndar col = mix(GOLD, RED, float(layer_ratio_norm[li])) width = max(2, round(2 + 11 * strength)) alpha = round(76 + 160 * strength) - draw_bezier(ed, ((sx + 162, sy), (1668, sy), (1748, out_anchor[1] + (sy - 760) * 0.18), out_anchor), fill=(*col, alpha), width=width) + draw_bezier( + ed, + ( + (sx + 162, sy), + (1668, sy), + (1748, out_anchor[1] + (sy - 760) * 0.18), + out_anchor, + ), + fill=(*col, alpha), + width=width, + ) edges = edges.filter(ImageFilter.GaussianBlur(0.18)) canvas = Image.alpha_composite(canvas.convert("RGBA"), edges).convert("RGB") @@ -383,12 +580,23 @@ def render(summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np.ndar legend_x, legend_y = 590, 1168 draw.text((legend_x, legend_y), "edge encoding", fill=INK, font=font(18, True)) - for i, (lab, val, col) in enumerate([("weak", 0.20, BLUE), ("medium", 0.55, GOLD), ("answer-specific", 0.95, RED)]): + for i, (lab, val, col) in enumerate( + [("weak", 0.20, BLUE), ("medium", 0.55, GOLD), ("answer-specific", 0.95, RED)] + ): yy = legend_y + 38 + i * 32 - draw.line((legend_x, yy, legend_x + 122, yy), fill=col, width=round(2 + 9 * val)) - draw.text((legend_x + 146, yy - 12), lab, fill=MUTED if i < 2 else INK, font=font(16)) + draw.line( + (legend_x, yy, legend_x + 122, yy), fill=col, width=round(2 + 9 * val) + ) + draw.text( + (legend_x + 146, yy - 12), lab, fill=MUTED if i < 2 else INK, font=font(16) + ) - draw.text((1158, 1164), "Data: heatmaps.npz answer_delta, random_delta, ratio. No schematic edges: every width/color is grouped from observed tensors.", fill=MUTED, font=font(17)) + draw.text( + (1158, 1164), + "Data: heatmaps.npz answer_delta, random_delta, ratio. No schematic edges: every width/color is grouped from observed tensors.", + fill=MUTED, + font=font(17), + ) out_dir.mkdir(parents=True, exist_ok=True) out_png = out_dir / "circuit.png" @@ -419,7 +627,14 @@ def main() -> None: out_dir = Path(args.out_dir) summary = json.loads((result_dir / "summary.json").read_text()) data = np.load(result_dir / "heatmaps.npz") - paths = render(summary, data["answer_delta"], data["random_delta"], data["ratio"], result_dir, out_dir) + paths = render( + summary, + data["answer_delta"], + data["random_delta"], + data["ratio"], + result_dir, + out_dir, + ) print(paths["png"]) diff --git a/packages/snapcompact/research/snapcompact_viz_city.py b/packages/snapcompact/research/snapcompact_viz_city.py index 5b34f9fb9..ce399f0ee 100644 --- a/packages/snapcompact/research/snapcompact_viz_city.py +++ b/packages/snapcompact/research/snapcompact_viz_city.py @@ -30,9 +30,15 @@ GRID = (39, 48, 83) def font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: candidates = [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/System/Library/Fonts/Supplemental/Helvetica Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Helvetica.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/System/Library/Fonts/Supplemental/Helvetica Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Helvetica.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ] for path in candidates: try: @@ -46,7 +52,9 @@ def clamp255(v: float) -> int: return max(0, min(255, int(round(v)))) -def mix(a: tuple[int, int, int], b: tuple[int, int, int], t: float) -> tuple[int, int, int]: +def mix( + a: tuple[int, int, int], b: tuple[int, int, int], t: float +) -> tuple[int, int, int]: t = max(0.0, min(1.0, t)) return tuple(clamp255(x + (y - x) * t) for x, y in zip(a, b)) @@ -55,12 +63,16 @@ def shade(c: tuple[int, int, int], factor: float) -> tuple[int, int, int]: return tuple(clamp255(x * factor) for x in c) -def iso(x: float, y: float, origin: tuple[float, float], tile_w: float, tile_h: float) -> tuple[float, float]: +def iso( + x: float, y: float, origin: tuple[float, float], tile_w: float, tile_h: float +) -> tuple[float, float]: ox, oy = origin return ox + (x - y) * tile_w * 0.5, oy + (x + y) * tile_h * 0.5 -def diamond(cx: float, cy: float, tile_w: float, tile_h: float) -> list[tuple[float, float]]: +def diamond( + cx: float, cy: float, tile_w: float, tile_h: float +) -> list[tuple[float, float]]: return [ (cx, cy - tile_h * 0.5), (cx + tile_w * 0.5, cy), @@ -69,14 +81,23 @@ def diamond(cx: float, cy: float, tile_w: float, tile_h: float) -> list[tuple[fl ] -def building_faces(cx: float, cy: float, h: float, tile_w: float, tile_h: float) -> tuple[list[tuple[float, float]], list[tuple[float, float]], list[tuple[float, float]]]: +def building_faces( + cx: float, cy: float, h: float, tile_w: float, tile_h: float +) -> tuple[ + list[tuple[float, float]], list[tuple[float, float]], list[tuple[float, float]] +]: top = diamond(cx, cy - h, tile_w, tile_h) right = [top[1], (cx + tile_w * 0.5, cy), (cx, cy + tile_h * 0.5), top[2]] left = [top[3], top[2], (cx, cy + tile_h * 0.5), (cx - tile_w * 0.5, cy)] return top, right, left -def draw_soft_line(draw: ImageDraw.ImageDraw, pts: Iterable[tuple[float, float]], fill: tuple[int, int, int], width: int = 1) -> None: +def draw_soft_line( + draw: ImageDraw.ImageDraw, + pts: Iterable[tuple[float, float]], + fill: tuple[int, int, int], + width: int = 1, +) -> None: draw.line([(int(x), int(y)) for x, y in pts], fill=fill, width=width) @@ -121,18 +142,37 @@ def draw_district( c = mix(base_color, high_color, max(intensity * 0.55, ratio_t * 0.85)) draw.polygon(left, fill=shade(c, 0.42)) draw.polygon(right, fill=shade(c, 0.62)) - draw.polygon(top, fill=mix(shade(c, 0.95), (255, 255, 255), intensity * 0.20)) + draw.polygon( + top, fill=mix(shade(c, 0.95), (255, 255, 255), intensity * 0.20) + ) if ratio_t > 0.80 or intensity > 0.90: - draw.line([(int(x), int(y)) for x, y in top + [top[0]]], fill=mix(c, (255, 255, 255), 0.25), width=1) + draw.line( + [(int(x), int(y)) for x, y in top + [top[0]]], + fill=mix(c, (255, 255, 255), 0.25), + width=1, + ) # District label plaque. x0, y0 = iso(-2, layers + 4, origin, tile_w, tile_h) x1, y1 = iso(54, layers + 4, origin, tile_w, tile_h) - draw.rounded_rectangle((x0 - 26, y0 + 18, x1 + 26, y1 + 64), radius=14, fill=(13, 18, 37), outline=shade(base_color, 0.75), width=2) - draw.text((x0 - 8, y0 + 27), label, font=font(26, True), fill=mix(base_color, high_color, 0.55)) + draw.rounded_rectangle( + (x0 - 26, y0 + 18, x1 + 26, y1 + 64), + radius=14, + fill=(13, 18, 37), + outline=shade(base_color, 0.75), + width=2, + ) + draw.text( + (x0 - 8, y0 + 27), + label, + font=font(26, True), + fill=mix(base_color, high_color, 0.55), + ) -def draw_legend(draw: ImageDraw.ImageDraw, summary: dict[str, object], scale: float) -> None: +def draw_legend( + draw: ImageDraw.ImageDraw, summary: dict[str, object], scale: float +) -> None: draw.text((88, 70), "Snapcompact Activation City", font=font(54, True), fill=INK) draw.text( (92, 136), @@ -140,26 +180,85 @@ def draw_legend(draw: ImageDraw.ImageDraw, summary: dict[str, object], scale: fl font=font(21), fill=MUTED, ) - question = str(summary.get("question", {}).get("q", "")) if isinstance(summary.get("question"), dict) else "" - answer = str(summary.get("question", {}).get("answer_text", "")) if isinstance(summary.get("question"), dict) else "" - draw.text((92, 173), f"Question: {question} Gold answer: {answer}", font=font(20), fill=(180, 190, 220)) + question = ( + str(summary.get("question", {}).get("q", "")) + if isinstance(summary.get("question"), dict) + else "" + ) + answer = ( + str(summary.get("question", {}).get("answer_text", "")) + if isinstance(summary.get("question"), dict) + else "" + ) + draw.text( + (92, 173), + f"Question: {question} Gold answer: {answer}", + font=font(20), + fill=(180, 190, 220), + ) ratio = float(summary.get("answer_over_random_delta", 0.0)) - draw.rounded_rectangle((1738, 72, 2286, 206), radius=24, fill=(12, 17, 36), outline=(50, 60, 99), width=2) - draw.text((1772, 96), "Answer-mask / random-mask mean delta", font=font(18), fill=MUTED) + draw.rounded_rectangle( + (1738, 72, 2286, 206), + radius=24, + fill=(12, 17, 36), + outline=(50, 60, 99), + width=2, + ) + draw.text( + (1772, 96), "Answer-mask / random-mask mean delta", font=font(18), fill=MUTED + ) draw.text((1772, 125), f"{ratio:.2f}×", font=font(52, True), fill=ANSWER_HI) - draw.text((1906, 149), f"common p98 height scale {scale:.1f}", font=font(17), fill=(176, 186, 218)) + draw.text( + (1906, 149), + f"common p98 height scale {scale:.1f}", + font=font(17), + fill=(176, 186, 218), + ) y = 1400 - draw.rounded_rectangle((88, y, 772, y + 96), radius=18, fill=(12, 17, 36), outline=(44, 54, 92), width=1) + draw.rounded_rectangle( + (88, y, 772, y + 96), + radius=18, + fill=(12, 17, 36), + outline=(44, 54, 92), + width=1, + ) draw.text((116, y + 18), "How to read it", font=font(22, True), fill=INK) - draw.text((116, y + 52), "Tall towers mark token/layer bins where masking changed hidden states most.", font=font(18), fill=MUTED) - draw.rounded_rectangle((836, y, 1520, y + 96), radius=18, fill=(12, 17, 36), outline=(44, 54, 92), width=1) + draw.text( + (116, y + 52), + "Tall towers mark token/layer bins where masking changed hidden states most.", + font=font(18), + fill=MUTED, + ) + draw.rounded_rectangle( + (836, y, 1520, y + 96), + radius=18, + fill=(12, 17, 36), + outline=(44, 54, 92), + width=1, + ) draw.text((864, y + 18), "Districts", font=font(22, True), fill=INK) - draw.text((864, y + 52), "Warm city = answer mask around “2003”; cool city = same-size random mask.", font=font(18), fill=MUTED) - draw.rounded_rectangle((1584, y, 2268, y + 96), radius=18, fill=(12, 17, 36), outline=(44, 54, 92), width=1) + draw.text( + (864, y + 52), + "Warm city = answer mask around “2003”; cool city = same-size random mask.", + font=font(18), + fill=MUTED, + ) + draw.rounded_rectangle( + (1584, y, 2268, y + 96), + radius=18, + fill=(12, 17, 36), + outline=(44, 54, 92), + width=1, + ) draw.text((1612, y + 18), "Color halos", font=font(22, True), fill=INK) - draw.text((1612, y + 52), "Bright caps emphasize bins with high answer/random activation ratio.", font=font(18), fill=MUTED) + draw.text( + (1612, y + 52), + "Bright caps emphasize bins with high answer/random activation ratio.", + font=font(18), + fill=MUTED, + ) def render() -> None: @@ -171,9 +270,14 @@ def render() -> None: random = np.asarray(heatmaps["random_binned"], dtype=np.float32) ratio = np.asarray(heatmaps["ratio_binned"], dtype=np.float32) if answer.shape != random.shape or answer.shape != ratio.shape: - raise ValueError(f"expected matching binned shapes, got {answer.shape}, {random.shape}, {ratio.shape}") + raise ValueError( + f"expected matching binned shapes, got {answer.shape}, {random.shape}, {ratio.shape}" + ) - scale = float(summary.get("common_delta_scale_p98") or np.percentile(np.concatenate([answer.ravel(), random.ravel()]), 98)) + scale = float( + summary.get("common_delta_scale_p98") + or np.percentile(np.concatenate([answer.ravel(), random.ravel()]), 98) + ) ratio_scale = float(summary.get("ratio_scale_p98") or np.percentile(ratio, 98)) w, h = 2400, 1600 @@ -185,18 +289,52 @@ def render() -> None: gd.ellipse((90, 250, 1090, 1260), fill=(255, 80, 95, 36)) gd.ellipse((1220, 250, 2250, 1260), fill=(62, 190, 255, 34)) gd.rectangle((0, 1240, w, h), fill=(2, 4, 11, 90)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(58))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(58)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) # Basemap plates. - draw.rounded_rectangle((52, 244, 1134, 1330), radius=42, fill=(9, 14, 30), outline=(42, 31, 55), width=2) - draw.rounded_rectangle((1234, 244, 2316, 1330), radius=42, fill=(8, 15, 31), outline=(24, 50, 70), width=2) + draw.rounded_rectangle( + (52, 244, 1134, 1330), + radius=42, + fill=(9, 14, 30), + outline=(42, 31, 55), + width=2, + ) + draw.rounded_rectangle( + (1234, 244, 2316, 1330), + radius=42, + fill=(8, 15, 31), + outline=(24, 50, 70), + width=2, + ) for y in range(312, 1300, 70): draw.line((70, y, 1116, y), fill=ROAD, width=1) draw.line((1252, y, 2298, y), fill=ROAD, width=1) - draw_district(draw, answer, ratio, (252.0, 870.0), ANSWER, ANSWER_HI, "answer-mask district", scale, ratio_scale) - draw_district(draw, random, ratio, (1434.0, 870.0), RANDOM, RANDOM_HI, "random-mask district", scale, ratio_scale) + draw_district( + draw, + answer, + ratio, + (252.0, 870.0), + ANSWER, + ANSWER_HI, + "answer-mask district", + scale, + ratio_scale, + ) + draw_district( + draw, + random, + ratio, + (1434.0, 870.0), + RANDOM, + RANDOM_HI, + "random-mask district", + scale, + ratio_scale, + ) draw_legend(draw, summary, scale) # Fine vignette frame. diff --git a/packages/snapcompact/research/snapcompact_viz_explainer.py b/packages/snapcompact/research/snapcompact_viz_explainer.py index bc89e27f3..ca63475d4 100644 --- a/packages/snapcompact/research/snapcompact_viz_explainer.py +++ b/packages/snapcompact/research/snapcompact_viz_explainer.py @@ -40,9 +40,15 @@ LINE = (37, 50, 64) def font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: names = [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/System/Library/Fonts/Supplemental/Helvetica Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Helvetica.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/System/Library/Fonts/Supplemental/Helvetica Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Helvetica.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ] for name in names: if name and Path(name).exists(): @@ -68,7 +74,9 @@ def lerp(a: int, b: int, t: float) -> int: return round(a + (b - a) * t) -def mix(a: tuple[int, int, int], b: tuple[int, int, int], t: float) -> tuple[int, int, int]: +def mix( + a: tuple[int, int, int], b: tuple[int, int, int], t: float +) -> tuple[int, int, int]: return tuple(lerp(a[i], b[i], t) for i in range(3)) @@ -94,19 +102,35 @@ def blue_color(t: float) -> tuple[int, int, int]: return mix((8, 13, 28), CYAN, t**0.75) -def paste_round(base: Image.Image, img: Image.Image, box: tuple[int, int, int, int], radius: int = 24) -> None: +def paste_round( + base: Image.Image, + img: Image.Image, + box: tuple[int, int, int, int], + radius: int = 24, +) -> None: x0, y0, x1, y1 = box img = img.convert("RGB") scale = min((x1 - x0) / img.width, (y1 - y0) / img.height) - resized = img.resize((max(1, round(img.width * scale)), max(1, round(img.height * scale))), Image.Resampling.LANCZOS) + resized = img.resize( + (max(1, round(img.width * scale)), max(1, round(img.height * scale))), + Image.Resampling.LANCZOS, + ) px = x0 + (x1 - x0 - resized.width) // 2 py = y0 + (y1 - y0 - resized.height) // 2 mask = Image.new("L", resized.size, 0) - ImageDraw.Draw(mask).rounded_rectangle((0, 0, resized.width - 1, resized.height - 1), radius=radius, fill=255) + ImageDraw.Draw(mask).rounded_rectangle( + (0, 0, resized.width - 1, resized.height - 1), radius=radius, fill=255 + ) base.paste(resized, (px, py), mask) -def draw_panel(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], title: str, subtitle: str, accent: tuple[int, int, int]) -> None: +def draw_panel( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + title: str, + subtitle: str, + accent: tuple[int, int, int], +) -> None: x0, y0, x1, y1 = box draw.rounded_rectangle(box, radius=30, fill=PANEL, outline=LINE, width=2) draw.rectangle((x0 + 28, y0 + 22, x0 + 84, y0 + 28), fill=accent) @@ -114,7 +138,13 @@ def draw_panel(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], title: draw.text((x0 + 28, y0 + 78), subtitle, fill=MUTED, font=F16) -def draw_arrow(draw: ImageDraw.ImageDraw, start: tuple[int, int], end: tuple[int, int], color: tuple[int, int, int], label: str) -> None: +def draw_arrow( + draw: ImageDraw.ImageDraw, + start: tuple[int, int], + end: tuple[int, int], + color: tuple[int, int, int], + label: str, +) -> None: sx, sy = start ex, ey = end draw.line((sx, sy, ex - 18, ey), fill=color, width=5) @@ -122,7 +152,12 @@ def draw_arrow(draw: ImageDraw.ImageDraw, start: tuple[int, int], end: tuple[int if not label: return tw = round(draw.textlength(label, font=F14)) - draw.rounded_rectangle((sx + 20, sy - 31, sx + 42 + tw, sy - 6), radius=12, fill=(9, 14, 24), outline=mix(color, LINE, 0.35)) + draw.rounded_rectangle( + (sx + 20, sy - 31, sx + 42 + tw, sy - 6), + radius=12, + fill=(9, 14, 24), + outline=mix(color, LINE, 0.35), + ) draw.text((sx + 31, sy - 29), label, fill=color, font=F14) @@ -137,10 +172,17 @@ def cell_box(summary: dict) -> tuple[int, int, int, int]: ch = 768 / rows r0, c0 = divmod(start, cols) r1, c1 = divmod(max(start, end - 1), cols) - return (math.floor(c0 * cw), math.floor(r0 * ch), math.ceil((c1 + 1) * cw), math.ceil((r1 + 1) * ch)) + return ( + math.floor(c0 * cw), + math.floor(r0 * ch), + math.ceil((c1 + 1) * cw), + math.ceil((r1 + 1) * ch), + ) -def answer_crop(img: Image.Image, summary: dict, pad_cells: int = 31) -> tuple[Image.Image, tuple[int, int, int, int]]: +def answer_crop( + img: Image.Image, summary: dict, pad_cells: int = 31 +) -> tuple[Image.Image, tuple[int, int, int, int]]: g = summary["geometry"] cols = int(g["cols"]) rows = int(g["rows"]) @@ -157,26 +199,49 @@ def answer_crop(img: Image.Image, summary: dict, pad_cells: int = 31) -> tuple[I y0 = max(0, math.floor((row - 5) * ch)) y1 = min(img.height, math.ceil((row + 6) * ch)) crop = img.crop((x0, y0, x1, y1)).convert("RGB") - local = (round(col0 * cw - x0), round(row * ch - y0), round(col1 * cw - x0), round((row + 1) * ch - y0)) + local = ( + round(col0 * cw - x0), + round(row * ch - y0), + round(col1 * cw - x0), + round((row + 1) * ch - y0), + ) return crop, local -def draw_crop_card(canvas: Image.Image, box: tuple[int, int, int, int], img: Image.Image, local_box: tuple[int, int, int, int], title: str, accent: tuple[int, int, int]) -> None: +def draw_crop_card( + canvas: Image.Image, + box: tuple[int, int, int, int], + img: Image.Image, + local_box: tuple[int, int, int, int], + title: str, + accent: tuple[int, int, int], +) -> None: draw = ImageDraw.Draw(canvas) x0, y0, x1, y1 = box draw.text((x0, y0 - 28), title, fill=accent, font=F16) - draw.rounded_rectangle(box, radius=18, fill=(236, 234, 219), outline=accent, width=3) + draw.rounded_rectangle( + box, radius=18, fill=(236, 234, 219), outline=accent, width=3 + ) pad = 14 scale = min((x1 - x0 - 2 * pad) / img.width, (y1 - y0 - 2 * pad) / img.height) - resized = img.resize((round(img.width * scale), round(img.height * scale)), Image.Resampling.NEAREST) + resized = img.resize( + (round(img.width * scale), round(img.height * scale)), Image.Resampling.NEAREST + ) px = x0 + (x1 - x0 - resized.width) // 2 py = y0 + (y1 - y0 - resized.height) // 2 canvas.paste(resized, (px, py)) bx = tuple(round(v * scale) for v in local_box) - draw.rounded_rectangle((px + bx[0] - 4, py + bx[1] - 4, px + bx[2] + 4, py + bx[3] + 4), radius=6, outline=accent, width=4) + draw.rounded_rectangle( + (px + bx[0] - 4, py + bx[1] - 4, px + bx[2] + 4, py + bx[3] + 4), + radius=6, + outline=accent, + width=4, + ) -def draw_heatmap(draw: ImageDraw.ImageDraw, arr: np.ndarray, box: tuple[int, int, int, int]) -> None: +def draw_heatmap( + draw: ImageDraw.ImageDraw, arr: np.ndarray, box: tuple[int, int, int, int] +) -> None: x0, y0, x1, y1 = box rows, cols = arr.shape cw = (x1 - x0) / cols @@ -190,11 +255,18 @@ def draw_heatmap(draw: ImageDraw.ImageDraw, arr: np.ndarray, box: tuple[int, int draw.rectangle((xa, ya, xb, yb), fill=heat_color(float(arr[r, c]))) for r in range(rows + 1): y = round(y0 + r * ch) - draw.line((x0, y, x1, y), fill=(0, 0, 0, 90) if False else (20, 26, 36), width=1) + draw.line( + (x0, y, x1, y), fill=(0, 0, 0, 90) if False else (20, 26, 36), width=1 + ) draw.rectangle(box, outline=(83, 101, 118), width=1) -def draw_tensor_ribbons(draw: ImageDraw.ImageDraw, answer: np.ndarray, random: np.ndarray, box: tuple[int, int, int, int]) -> None: +def draw_tensor_ribbons( + draw: ImageDraw.ImageDraw, + answer: np.ndarray, + random: np.ndarray, + box: tuple[int, int, int, int], +) -> None: x0, y0, x1, y1 = box rows, cols = answer.shape lane_h = (y1 - y0) / rows @@ -215,7 +287,9 @@ def draw_tensor_ribbons(draw: ImageDraw.ImageDraw, answer: np.ndarray, random: n draw.rectangle(box, outline=(91, 106, 122), width=1) -def draw_token_grid(draw: ImageDraw.ImageDraw, ratio: np.ndarray, box: tuple[int, int, int, int]) -> None: +def draw_token_grid( + draw: ImageDraw.ImageDraw, ratio: np.ndarray, box: tuple[int, int, int, int] +) -> None: x0, y0, x1, y1 = box grid = ratio.mean(axis=0).reshape(27, 27) q98 = float(np.quantile(grid, 0.98)) or 1.0 @@ -229,15 +303,28 @@ def draw_token_grid(draw: ImageDraw.ImageDraw, ratio: np.ndarray, box: tuple[int ya = round(oy + r * cell) xb = round(ox + (c + 1) * cell - 1) yb = round(oy + (r + 1) * cell - 1) - draw.rounded_rectangle((xa, ya, xb, yb), radius=3, fill=blue_color(float(norm[r, c]))) + draw.rounded_rectangle( + (xa, ya, xb, yb), radius=3, fill=blue_color(float(norm[r, c])) + ) top = np.unravel_index(np.argsort(grid, axis=None)[-6:], grid.shape) for r, c in zip(top[0], top[1]): xa = round(ox + c * cell) ya = round(oy + r * cell) - draw.rounded_rectangle((xa - 2, ya - 2, round(xa + cell + 1), round(ya + cell + 1)), radius=4, outline=AMBER, width=2) + draw.rounded_rectangle( + (xa - 2, ya - 2, round(xa + cell + 1), round(ya + cell + 1)), + radius=4, + outline=AMBER, + width=2, + ) -def polyline(draw: ImageDraw.ImageDraw, values: Iterable[float], box: tuple[int, int, int, int], color: tuple[int, int, int], width: int = 4) -> None: +def polyline( + draw: ImageDraw.ImageDraw, + values: Iterable[float], + box: tuple[int, int, int, int], + color: tuple[int, int, int], + width: int = 4, +) -> None: vals = list(values) x0, y0, x1, y1 = box lo = min(vals) @@ -254,15 +341,26 @@ def polyline(draw: ImageDraw.ImageDraw, values: Iterable[float], box: tuple[int, draw.ellipse((x - 3, y - 3, x + 3, y + 3), fill=color) -def metric(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], label: str, value: str, sub: str, accent: tuple[int, int, int]) -> None: +def metric( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + label: str, + value: str, + sub: str, + accent: tuple[int, int, int], +) -> None: x0, y0, x1, y1 = box - draw.rounded_rectangle(box, radius=20, fill=PANEL_2, outline=mix(accent, LINE, 0.35), width=2) + draw.rounded_rectangle( + box, radius=20, fill=PANEL_2, outline=mix(accent, LINE, 0.35), width=2 + ) draw.text((x0 + 18, y0 + 16), label, fill=MUTED, font=F14) draw.text((x0 + 18, y0 + 41), value, fill=accent, font=F38) draw.text((x0 + 18, y1 - 32), sub, fill=INK, font=F14) -def wrap_text(draw: ImageDraw.ImageDraw, text: str, max_width: int, fnt: ImageFont.ImageFont) -> list[str]: +def wrap_text( + draw: ImageDraw.ImageDraw, text: str, max_width: int, fnt: ImageFont.ImageFont +) -> list[str]: words = text.split() lines: list[str] = [] cur = "" @@ -279,7 +377,9 @@ def wrap_text(draw: ImageDraw.ImageDraw, text: str, max_width: int, fnt: ImageFo return lines -def save_source_metrics(out_dir: Path, summary: dict, arrays: dict[str, np.ndarray]) -> None: +def save_source_metrics( + out_dir: Path, summary: dict, arrays: dict[str, np.ndarray] +) -> None: ratio = arrays["ratio"] answer = arrays["answer_delta"] random = arrays["random_delta"] @@ -300,7 +400,9 @@ def save_source_metrics(out_dir: Path, summary: dict, arrays: dict[str, np.ndarr "mean_ratio_by_layer": [float(x) for x in ratio.mean(axis=1)], } out_dir.mkdir(parents=True, exist_ok=True) - (out_dir / "explainer_metrics.json").write_text(json.dumps(metrics, indent=2) + "\n") + (out_dir / "explainer_metrics.json").write_text( + json.dumps(metrics, indent=2) + "\n" + ) def render(data_dir: Path, out_dir: Path) -> Path: @@ -324,7 +426,9 @@ def render(data_dir: Path, out_dir: Path) -> Path: gd.ellipse((-300, -210, 980, 760), fill=(78, 219, 255, 35)) gd.ellipse((690, 200, 1850, 1370), fill=(255, 82, 65, 32)) gd.ellipse((1550, -120, 2660, 980), fill=(255, 197, 78, 24)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(90))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(90)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) draw.text((72, 52), "SNAPCOMPACT ACTIVATION EXPLAINER", fill=AMBER, font=F22) @@ -338,9 +442,15 @@ def render(data_dir: Path, out_dir: Path) -> Path: p2 = (630, 240, 1143, 1178) p3 = (1188, 240, 1770, 1178) p4 = (1815, 240, 2328, 1178) - draw_panel(draw, p1, "1 · input bitmap", "Rendered context before intervention", CYAN) - draw_panel(draw, p2, "2 · mask intervention", "Only the true answer cells are blanked", RED) - draw_panel(draw, p3, "3 · hidden-state tensor", "Layer × image-token response", VIOLET) + draw_panel( + draw, p1, "1 · input bitmap", "Rendered context before intervention", CYAN + ) + draw_panel( + draw, p2, "2 · mask intervention", "Only the true answer cells are blanked", RED + ) + draw_panel( + draw, p3, "3 · hidden-state tensor", "Layer × image-token response", VIOLET + ) draw_panel(draw, p4, "4 · interpretation", "Where the answer mattered most", AMBER) draw_arrow(draw, (585, 715), (630, 715), CYAN, "") draw_arrow(draw, (1143, 715), (1188, 715), RED, "") @@ -351,44 +461,144 @@ def render(data_dir: Path, out_dir: Path) -> Path: full_box = cell_box(summary) scale = 440 / 768 ox, oy = 108, 352 - draw.rounded_rectangle((ox + round(full_box[0] * scale), oy + round(full_box[1] * scale), ox + round(full_box[2] * scale), oy + round(full_box[3] * scale)), radius=5, outline=AMBER, width=4) + draw.rounded_rectangle( + ( + ox + round(full_box[0] * scale), + oy + round(full_box[1] * scale), + ox + round(full_box[2] * scale), + oy + round(full_box[3] * scale), + ), + radius=5, + outline=AMBER, + width=4, + ) ocrop, local = answer_crop(original, summary) mcrop, mlocal = answer_crop(answer_mask, summary) - draw_crop_card(canvas, (108, 885, 548, 1042), ocrop, local, "magnified answer glyphs", AMBER) - draw.text((108, 1083), "The OCR input is a fixed bitmap. The answer span", fill=MUTED, font=F16) - draw.text((108, 1108), f"occupies character cells {summary['question']['answer_start']}–{summary['question']['answer_end'] - 1}.", fill=MUTED, font=F16) + draw_crop_card( + canvas, (108, 885, 548, 1042), ocrop, local, "magnified answer glyphs", AMBER + ) + draw.text( + (108, 1083), + "The OCR input is a fixed bitmap. The answer span", + fill=MUTED, + font=F16, + ) + draw.text( + (108, 1108), + f"occupies character cells {summary['question']['answer_start']}–{summary['question']['answer_end'] - 1}.", + fill=MUTED, + font=F16, + ) paste_round(canvas, answer_mask, (666, 352, 1106, 792), 24) - draw.rounded_rectangle((666 + round(full_box[0] * scale), 352 + round(full_box[1] * scale), 666 + round(full_box[2] * scale), 352 + round(full_box[3] * scale)), radius=5, outline=RED, width=4) - draw_crop_card(canvas, (666, 885, 1106, 1042), mcrop, mlocal, "same crop after masking", RED) - draw.text((666, 1083), "Same prompt, same rendered page. Difference:", fill=MUTED, font=F16) - draw.text((666, 1108), "the four answer glyphs are removed before inference.", fill=MUTED, font=F16) + draw.rounded_rectangle( + ( + 666 + round(full_box[0] * scale), + 352 + round(full_box[1] * scale), + 666 + round(full_box[2] * scale), + 352 + round(full_box[3] * scale), + ), + radius=5, + outline=RED, + width=4, + ) + draw_crop_card( + canvas, (666, 885, 1106, 1042), mcrop, mlocal, "same crop after masking", RED + ) + draw.text( + (666, 1083), + "Same prompt, same rendered page. Difference:", + fill=MUTED, + font=F16, + ) + draw.text( + (666, 1108), + "the four answer glyphs are removed before inference.", + fill=MUTED, + font=F16, + ) # Tensor panel. draw.text((1226, 332), "answer-mask delta", fill=RED, font=F18) draw.text((1600, 332), "random control mixed in green", fill=GREEN, font=F14) - draw_tensor_ribbons(draw, arrays["answer_norm"], arrays["random_norm"], (1240, 372, 1718, 665)) + draw_tensor_ribbons( + draw, arrays["answer_norm"], arrays["random_norm"], (1240, 372, 1718, 665) + ) draw.text((1238, 686), "answer / random ratio", fill=AMBER, font=F18) - draw.text((1238, 712), "bright = answer-region deletion moves hidden states more than an equal random mask", fill=MUTED, font=F14) + draw.text( + (1238, 712), + "bright = answer-region deletion moves hidden states more than an equal random mask", + fill=MUTED, + font=F14, + ) draw_heatmap(draw, arrays["ratio_norm"], (1240, 748, 1718, 1000)) max_layer = int(summary["max_ratio_layer"]) - draw.line((1240, 748 + round((max_layer + 0.5) * 252 / 19), 1718, 748 + round((max_layer + 0.5) * 252 / 19)), fill=AMBER, width=3) + draw.line( + ( + 1240, + 748 + round((max_layer + 0.5) * 252 / 19), + 1718, + 748 + round((max_layer + 0.5) * 252 / 19), + ), + fill=AMBER, + width=3, + ) for i in range(240): draw.rectangle((1240 + i, 1046, 1241 + i, 1063), fill=heat_color(i / 239)) draw.text((1240, 1022), "low", fill=MUTED, font=F12) draw.text((1446, 1022), "high", fill=MUTED, font=F12) - draw.text((1240, 1094), f"{summary['layers']} decoder layers × {summary['image_tokens']} image tokens", fill=INK, font=F20) - draw.text((1240, 1124), "Each cell uses the saved heatmaps.npz tensor values.", fill=MUTED, font=F16) + draw.text( + (1240, 1094), + f"{summary['layers']} decoder layers × {summary['image_tokens']} image tokens", + fill=INK, + font=F20, + ) + draw.text( + (1240, 1124), + "Each cell uses the saved heatmaps.npz tensor values.", + fill=MUTED, + font=F16, + ) # Interpretation panel. - metric(draw, (1850, 344, 2075, 478), "mean delta ratio", f"{summary['answer_over_random_delta']:.2f}×", "answer mask vs control", AMBER) - metric(draw, (2086, 344, 2293, 478), "strongest layer", f"L{summary['max_ratio_layer']}", "mean ratio peak", VIOLET) - metric(draw, (1850, 500, 2075, 634), "answer delta", f"{summary['answer_delta_mean']:.2f}", "mean ||Δh||", RED) - metric(draw, (2086, 500, 2293, 634), "control delta", f"{summary['random_delta_mean']:.2f}", "mean ||Δh||", GREEN) + metric( + draw, + (1850, 344, 2075, 478), + "mean delta ratio", + f"{summary['answer_over_random_delta']:.2f}×", + "answer mask vs control", + AMBER, + ) + metric( + draw, + (2086, 344, 2293, 478), + "strongest layer", + f"L{summary['max_ratio_layer']}", + "mean ratio peak", + VIOLET, + ) + metric( + draw, + (1850, 500, 2075, 634), + "answer delta", + f"{summary['answer_delta_mean']:.2f}", + "mean ||Δh||", + RED, + ) + metric( + draw, + (2086, 500, 2293, 634), + "control delta", + f"{summary['random_delta_mean']:.2f}", + "mean ||Δh||", + GREEN, + ) draw.text((1852, 684), "layer sensitivity curve", fill=INK, font=F20) curve_box = (1862, 725, 2290, 858) - draw.rounded_rectangle((1850, 704, 2304, 884), radius=20, fill=PANEL_2, outline=LINE, width=2) + draw.rounded_rectangle( + (1850, 704, 2304, 884), radius=20, fill=PANEL_2, outline=LINE, width=2 + ) for i in range(5): y = curve_box[1] + i * (curve_box[3] - curve_box[1]) / 4 draw.line((curve_box[0], round(y), curve_box[2], round(y)), fill=(31, 42, 54)) @@ -397,14 +607,23 @@ def render(data_dir: Path, out_dir: Path) -> Path: draw.text((2262, 862), f"L{summary['layers'] - 1}", fill=MUTED, font=F12) draw.text((1852, 927), "image-token sensitivity field", fill=INK, font=F20) - draw.rounded_rectangle((1850, 955, 2067, 1150), radius=20, fill=PANEL_2, outline=LINE, width=2) + draw.rounded_rectangle( + (1850, 955, 2067, 1150), radius=20, fill=PANEL_2, outline=LINE, width=2 + ) draw_token_grid(draw, arrays["ratio"], (1868, 970, 2049, 1132)) explanation = "Answer deletion creates a high-ratio band in early layers; later layers diffuse it into surrounding context." for i, line in enumerate(wrap_text(draw, explanation, 195, F14)): draw.text((2092, 968 + i * 24), line, fill=INK if i == 0 else MUTED, font=F14) draw.text((2092, 1090), "Interpretation:", fill=AMBER, font=F16) - draw.text((2092, 1118), "the answer glyphs are not just OCR text;", fill=MUTED, font=F14) - draw.text((2092, 1142), "they perturb the multimodal residual stream.", fill=MUTED, font=F14) + draw.text( + (2092, 1118), "the answer glyphs are not just OCR text;", fill=MUTED, font=F14 + ) + draw.text( + (2092, 1142), + "they perturb the multimodal residual stream.", + fill=MUTED, + font=F14, + ) # Footer with provenance. footer = (72, 1228, 2328, 1422) @@ -413,12 +632,26 @@ def render(data_dir: Path, out_dir: Path) -> Path: bullets = [ (CYAN, "Input bitmap", "is the rendered evidence page passed to PaddleOCR-VL."), (RED, "Mask intervention", "removes only the gold answer span: 2003."), - (VIOLET, "Hidden-state tensor", "plots ||hidden(original) − hidden(masked)|| over saved layer/token arrays."), - (AMBER, "Interpretation", "compares that scar to an equal-size random mask: 2.52× stronger on average."), + ( + VIOLET, + "Hidden-state tensor", + "plots ||hidden(original) − hidden(masked)|| over saved layer/token arrays.", + ), + ( + AMBER, + "Interpretation", + "compares that scar to an equal-size random mask: 2.52× stronger on average.", + ), ] x = 108 for color, head, text in bullets: - draw.rounded_rectangle((x, 1320, x + 500, 1384), radius=18, fill=PANEL_2, outline=mix(color, LINE, 0.35), width=2) + draw.rounded_rectangle( + (x, 1320, x + 500, 1384), + radius=18, + fill=PANEL_2, + outline=mix(color, LINE, 0.35), + width=2, + ) draw.ellipse((x + 18, 1343, x + 36, 1361), fill=color) draw.text((x + 50, 1330), head, fill=color, font=F16) draw.text((x + 50, 1355), text, fill=MUTED, font=F14) diff --git a/packages/snapcompact/research/snapcompact_viz_glass_stack.py b/packages/snapcompact/research/snapcompact_viz_glass_stack.py index 29b206203..4c77fb3d7 100644 --- a/packages/snapcompact/research/snapcompact_viz_glass_stack.py +++ b/packages/snapcompact/research/snapcompact_viz_glass_stack.py @@ -28,9 +28,13 @@ PANEL = (9, 16, 26) def font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: candidates = [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", "/System/Library/Fonts/Helvetica.ttc", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ] for candidate in candidates: if candidate and Path(candidate).exists(): @@ -38,7 +42,9 @@ def font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.Im return ImageFont.load_default() -def mix(a: tuple[int, int, int], b: tuple[int, int, int], t: float) -> tuple[int, int, int]: +def mix( + a: tuple[int, int, int], b: tuple[int, int, int], t: float +) -> tuple[int, int, int]: t = max(0.0, min(1.0, t)) return tuple(round(a[i] + (b[i] - a[i]) * t) for i in range(3)) @@ -59,7 +65,11 @@ def glass_heat(t: float) -> tuple[int, int, int]: return stops[-1][1] -def plane_corners(layer: int) -> tuple[tuple[float, float], tuple[float, float], tuple[float, float], tuple[float, float]]: +def plane_corners( + layer: int, +) -> tuple[ + tuple[float, float], tuple[float, float], tuple[float, float], tuple[float, float] +]: x = 225 + layer * 25.0 y = 920 - layer * 34.0 width = 930.0 @@ -67,7 +77,16 @@ def plane_corners(layer: int) -> tuple[tuple[float, float], tuple[float, float], return (x, y), (x + width, y), (x + width + dx, y + dy), (x + dx, y + dy) -def bilerp(corners: tuple[tuple[float, float], tuple[float, float], tuple[float, float], tuple[float, float]], u: float, v: float) -> tuple[float, float]: +def bilerp( + corners: tuple[ + tuple[float, float], + tuple[float, float], + tuple[float, float], + tuple[float, float], + ], + u: float, + v: float, +) -> tuple[float, float]: fl, fr, br, bl = corners ax = fl[0] + (fr[0] - fl[0]) * u ay = fl[1] + (fr[1] - fl[1]) * u @@ -80,11 +99,18 @@ def poly(points: Iterable[tuple[float, float]]) -> list[tuple[int, int]]: return [(round(x), round(y)) for x, y in points] -def select_scars(ratio_norm: np.ndarray, answer_binned: np.ndarray, random_binned: np.ndarray, count: int = 7) -> list[int]: +def select_scars( + ratio_norm: np.ndarray, + answer_binned: np.ndarray, + random_binned: np.ndarray, + count: int = 7, +) -> list[int]: advantage = np.maximum(answer_binned - random_binned, 0.0) if float(advantage.max(initial=0.0)) > 0: advantage = advantage / float(np.quantile(advantage, 0.985)) - score = ratio_norm.mean(axis=0) * 0.68 + np.clip(advantage, 0, 1).mean(axis=0) * 0.32 + score = ( + ratio_norm.mean(axis=0) * 0.68 + np.clip(advantage, 0, 1).mean(axis=0) * 0.32 + ) order = np.argsort(score)[::-1] chosen: list[int] = [] for idx in order: @@ -129,29 +155,66 @@ def draw_plane(canvas: Image.Image, values: np.ndarray, layer: int) -> None: rgb = glass_heat(shade) alpha = round(26 + 96 * math.pow(shade, 0.82)) draw.polygon( - poly((bilerp(corners, u0, 0.03), bilerp(corners, u1, 0.03), bilerp(corners, u1, 0.97), bilerp(corners, u0, 0.97))), + poly( + ( + bilerp(corners, u0, 0.03), + bilerp(corners, u1, 0.03), + bilerp(corners, u1, 0.97), + bilerp(corners, u0, 0.97), + ) + ), fill=(*rgb, alpha), ) for u in np.linspace(0, 1, 13): - draw.line(poly((bilerp(corners, float(u), 0), bilerp(corners, float(u), 1))), fill=(190, 242, 255, 28), width=1) + draw.line( + poly((bilerp(corners, float(u), 0), bilerp(corners, float(u), 1))), + fill=(190, 242, 255, 28), + width=1, + ) for v in np.linspace(0, 1, 5): - draw.line(poly((bilerp(corners, 0, float(v)), bilerp(corners, 1, float(v)))), fill=(190, 242, 255, 24), width=1) - draw.line(poly((corners[0], corners[1], corners[2], corners[3], corners[0])), fill=(174, 241, 255, 70), width=2) + draw.line( + poly((bilerp(corners, 0, float(v)), bilerp(corners, 1, float(v)))), + fill=(190, 242, 255, 24), + width=1, + ) + draw.line( + poly((corners[0], corners[1], corners[2], corners[3], corners[0])), + fill=(174, 241, 255, 70), + width=2, + ) if layer in (0, 6, 12, 18): x, y = corners[0] - draw.text((round(x - 64), round(y - 10)), f"L{layer:02d}", fill=(178, 226, 239, 150), font=font(16, True)) + draw.text( + (round(x - 64), round(y - 10)), + f"L{layer:02d}", + fill=(178, 226, 239, 150), + font=font(16, True), + ) canvas.alpha_composite(overlay) -def draw_scars(canvas: Image.Image, scar_bins: list[int], ratio_norm: np.ndarray) -> None: +def draw_scars( + canvas: Image.Image, scar_bins: list[int], ratio_norm: np.ndarray +) -> None: glow = Image.new("RGBA", canvas.size, (0, 0, 0, 0)) gd = ImageDraw.Draw(glow, "RGBA") cols = ratio_norm.shape[1] - scar_colors = [(255, 238, 164), (255, 102, 132), (87, 236, 255), (255, 190, 80), (205, 111, 255), (255, 255, 255), (72, 255, 190)] + scar_colors = [ + (255, 238, 164), + (255, 102, 132), + (87, 236, 255), + (255, 190, 80), + (205, 111, 255), + (255, 255, 255), + (72, 255, 190), + ] for n, c in enumerate(scar_bins): u = (c + 0.5) / cols - pts = [bilerp(plane_corners(layer), u, 0.46) for layer in range(ratio_norm.shape[0])] + pts = [ + bilerp(plane_corners(layer), u, 0.46) + for layer in range(ratio_norm.shape[0]) + ] color = scar_colors[n % len(scar_colors)] gd.line(poly(pts), fill=(*color, 120), width=12) for layer, pt in enumerate(pts): @@ -163,42 +226,107 @@ def draw_scars(canvas: Image.Image, scar_bins: list[int], ratio_norm: np.ndarray draw = ImageDraw.Draw(canvas, "RGBA") for n, c in enumerate(scar_bins): u = (c + 0.5) / cols - pts = [bilerp(plane_corners(layer), u, 0.46) for layer in range(ratio_norm.shape[0])] + pts = [ + bilerp(plane_corners(layer), u, 0.46) + for layer in range(ratio_norm.shape[0]) + ] color = scar_colors[n % len(scar_colors)] draw.line(poly(pts), fill=(*color, 235), width=3) top = pts[-1] - draw.text((round(top[0] + 10), round(top[1] - 14)), f"bin {c}", fill=(*color, 220), font=font(13, True)) + draw.text( + (round(top[0] + 10), round(top[1] - 14)), + f"bin {c}", + fill=(*color, 220), + font=font(13, True), + ) for layer, pt in enumerate(pts): r = 2.0 + 4.5 * float(ratio_norm[layer, c]) x, y = pt - draw.ellipse((x - r, y - r, x + r, y + r), fill=(255, 255, 230, 225), outline=(*color, 255), width=1) + draw.ellipse( + (x - r, y - r, x + r, y + r), + fill=(255, 255, 230, 225), + outline=(*color, 255), + width=1, + ) -def draw_labels(canvas: Image.Image, summary: dict, scar_bins: list[int], ratio_binned: np.ndarray) -> None: +def draw_labels( + canvas: Image.Image, summary: dict, scar_bins: list[int], ratio_binned: np.ndarray +) -> None: draw = ImageDraw.Draw(canvas, "RGBA") q = summary["question"]["q"] answer = summary["question"]["answer_text"] draw.text((70, 54), "SNAPCOMPACT GLASS STACK", fill=GOLD, font=font(22, True)) - draw.text((70, 88), "Answer-mask scars through decoder depth", fill=INK, font=font(54, True)) - draw.text((73, 154), f"Question: {q} · gold answer: {answer}", fill=MUTED, font=font(22)) + draw.text( + (70, 88), + "Answer-mask scars through decoder depth", + fill=INK, + font=font(54, True), + ) + draw.text( + (73, 154), + f"Question: {q} · gold answer: {answer}", + fill=MUTED, + font=font(22), + ) x0, y0, x1, y1 = 70, 960, 770, 1110 - draw.rounded_rectangle((x0, y0, x1, y1), radius=24, fill=(7, 13, 22, 205), outline=(115, 217, 255, 72), width=1) + draw.rounded_rectangle( + (x0, y0, x1, y1), + radius=24, + fill=(7, 13, 22, 205), + outline=(115, 217, 255, 72), + width=1, + ) ratio = summary["answer_over_random_delta"] draw.text((x0 + 26, y0 + 22), f"{ratio:.2f}×", fill=GOLD, font=font(48, True)) - draw.text((x0 + 170, y0 + 31), "mean answer-mask / random-mask delta", fill=INK, font=font(22, True)) - draw.text((x0 + 28, y0 + 86), f"{summary['layers']} semi-transparent decoder planes · {summary['image_tokens']} image tokens binned to {ratio_binned.shape[1]} columns", fill=MUTED, font=font(18)) + draw.text( + (x0 + 170, y0 + 31), + "mean answer-mask / random-mask delta", + fill=INK, + font=font(22, True), + ) + draw.text( + (x0 + 28, y0 + 86), + f"{summary['layers']} semi-transparent decoder planes · {summary['image_tokens']} image tokens binned to {ratio_binned.shape[1]} columns", + fill=MUTED, + font=font(18), + ) lx0, ly0 = 1240, 930 - draw.rounded_rectangle((lx0, ly0, lx0 + 475, ly0 + 182), radius=24, fill=(7, 13, 22, 210), outline=(115, 217, 255, 70), width=1) + draw.rounded_rectangle( + (lx0, ly0, lx0 + 475, ly0 + 182), + radius=24, + fill=(7, 13, 22, 210), + outline=(115, 217, 255, 70), + width=1, + ) draw.text((lx0 + 24, ly0 + 22), "encoding", fill=INK, font=font(25, True)) - draw.text((lx0 + 24, ly0 + 61), "plane color = answer/random ratio", fill=MUTED, font=font(18)) - draw.text((lx0 + 24, ly0 + 92), "vertical scar = high-ratio token bin", fill=MUTED, font=font(18)) - draw.text((lx0 + 24, ly0 + 124), "selected bins: " + ", ".join(map(str, scar_bins)), fill=(203, 231, 240), font=font(17)) + draw.text( + (lx0 + 24, ly0 + 61), + "plane color = answer/random ratio", + fill=MUTED, + font=font(18), + ) + draw.text( + (lx0 + 24, ly0 + 92), + "vertical scar = high-ratio token bin", + fill=MUTED, + font=font(18), + ) + draw.text( + (lx0 + 24, ly0 + 124), + "selected bins: " + ", ".join(map(str, scar_bins)), + fill=(203, 231, 240), + font=font(17), + ) # Color ramp. for i in range(220): - draw.rectangle((lx0 + 230 + i, ly0 + 30, lx0 + 231 + i, ly0 + 49), fill=(*glass_heat(i / 219), 255)) + draw.rectangle( + (lx0 + 230 + i, ly0 + 30, lx0 + 231 + i, ly0 + 49), + fill=(*glass_heat(i / 219), 255), + ) draw.text((lx0 + 230, ly0 + 54), "low", fill=MUTED, font=font(13)) draw.text((lx0 + 417, ly0 + 54), "high", fill=MUTED, font=font(13)) @@ -223,7 +351,9 @@ def main() -> None: draw_scars(canvas, scar_bins, ratio_norm) draw_labels(canvas, summary, scar_bins, ratio_binned) - ImageDraw.Draw(canvas).rounded_rectangle((42, 36, 1760, 1142), radius=38, outline=(128, 225, 255, 44), width=2) + ImageDraw.Draw(canvas).rounded_rectangle( + (42, 36, 1760, 1142), radius=38, outline=(128, 225, 255, 44), width=2 + ) out_path = OUT_DIR / "glass-stack.png" canvas.convert("RGB").save(out_path, quality=95) diff --git a/packages/snapcompact/research/snapcompact_viz_glyph_matrix.py b/packages/snapcompact/research/snapcompact_viz_glyph_matrix.py index 77dd63f19..05973c502 100644 --- a/packages/snapcompact/research/snapcompact_viz_glyph_matrix.py +++ b/packages/snapcompact/research/snapcompact_viz_glyph_matrix.py @@ -42,9 +42,13 @@ PALETTE = { def font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: candidates = [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", "/System/Library/Fonts/Supplemental/Helvetica.ttc", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ] for path in candidates: if path and Path(path).exists(): @@ -52,7 +56,9 @@ def font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.Im return ImageFont.load_default() -def mix(a: tuple[int, int, int], b: tuple[int, int, int], t: float) -> tuple[int, int, int]: +def mix( + a: tuple[int, int, int], b: tuple[int, int, int], t: float +) -> tuple[int, int, int]: t = max(0.0, min(1.0, t)) return tuple(round(a[i] + (b[i] - a[i]) * t) for i in range(3)) @@ -80,12 +86,23 @@ def quantile_norm(values: np.ndarray, q: float = 0.98) -> np.ndarray: return np.clip(values / scale, 0.0, 1.0) -def rounded_panel(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], title: str | None = None, subtitle: str | None = None) -> None: - draw.rounded_rectangle(box, radius=24, fill=PALETTE["panel"], outline=(34, 48, 61), width=1) +def rounded_panel( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + title: str | None = None, + subtitle: str | None = None, +) -> None: + draw.rounded_rectangle( + box, radius=24, fill=PALETTE["panel"], outline=(34, 48, 61), width=1 + ) if title: - draw.text((box[0] + 24, box[1] + 18), title, fill=PALETTE["ink"], font=font(28, True)) + draw.text( + (box[0] + 24, box[1] + 18), title, fill=PALETTE["ink"], font=font(28, True) + ) if subtitle: - draw.text((box[0] + 24, box[1] + 54), subtitle, fill=PALETTE["muted"], font=font(16)) + draw.text( + (box[0] + 24, box[1] + 54), subtitle, fill=PALETTE["muted"], font=font(16) + ) def token_boxes(side: int, grid: int) -> list[tuple[int, int, int, int]]: @@ -100,7 +117,9 @@ def token_boxes(side: int, grid: int) -> list[tuple[int, int, int, int]]: return boxes -def intersect_area(a: tuple[float, float, float, float], b: tuple[float, float, float, float]) -> float: +def intersect_area( + a: tuple[float, float, float, float], b: tuple[float, float, float, float] +) -> float: x0 = max(a[0], b[0]) y0 = max(a[1], b[1]) x1 = min(a[2], b[2]) @@ -108,7 +127,9 @@ def intersect_area(a: tuple[float, float, float, float], b: tuple[float, float, return max(0.0, x1 - x0) * max(0.0, y1 - y0) -def answer_bbox(start: int, end: int, cols: int, adv: int, pitch: int) -> tuple[int, int, int, int]: +def answer_bbox( + start: int, end: int, cols: int, adv: int, pitch: int +) -> tuple[int, int, int, int]: row0, col0 = divmod(start, cols) row1, col1 = divmod(max(start, end - 1), cols) x0 = max(0, col0 * adv) @@ -118,7 +139,16 @@ def answer_bbox(start: int, end: int, cols: int, adv: int, pitch: int) -> tuple[ return x0, y0, x1, y1 -def draw_text_wrapped(draw: ImageDraw.ImageDraw, xy: tuple[int, int], text: str, width: int, fill: tuple[int, int, int], size: int, bold: bool = False, line_gap: int = 4) -> int: +def draw_text_wrapped( + draw: ImageDraw.ImageDraw, + xy: tuple[int, int], + text: str, + width: int, + fill: tuple[int, int, int], + size: int, + bold: bool = False, + line_gap: int = 4, +) -> int: words = text.split() lines: list[str] = [] current = "" @@ -143,7 +173,9 @@ def paste_shadowed(canvas: Image.Image, img: Image.Image, xy: tuple[int, int]) - shadow = Image.new("RGBA", img.size, (0, 0, 0, 0)) alpha = Image.new("L", img.size, 180) shadow.putalpha(alpha) - canvas.alpha_composite(shadow.filter(ImageFilter.GaussianBlur(12)), (xy[0] + 8, xy[1] + 10)) + canvas.alpha_composite( + shadow.filter(ImageFilter.GaussianBlur(12)), (xy[0] + 8, xy[1] + 10) + ) canvas.alpha_composite(img, xy) @@ -175,7 +207,9 @@ def draw_activation_overlay( for idx in top_tokens: box = boxes[int(idx)] - draw.rounded_rectangle(box, radius=3, outline=PALETTE["amber"] + (235,), width=3) + draw.rounded_rectangle( + box, radius=3, outline=PALETTE["amber"] + (235,), width=3 + ) for idx in answer_tokens: box = boxes[int(idx)] draw.rounded_rectangle(box, radius=4, outline=PALETTE["cyan"] + (245,), width=4) @@ -183,13 +217,34 @@ def draw_activation_overlay( glow = Image.new("RGBA", composite.size, (0, 0, 0, 0)) gd = ImageDraw.Draw(glow) for w, a in ((16, 46), (9, 80), (4, 235)): - gd.rounded_rectangle((bbox[0] - 8, bbox[1] - 7, bbox[2] + 8, bbox[3] + 8), radius=8, outline=PALETTE["red"] + (a,), width=w) - composite = Image.alpha_composite(composite, glow.filter(ImageFilter.GaussianBlur(4))) + gd.rounded_rectangle( + (bbox[0] - 8, bbox[1] - 7, bbox[2] + 8, bbox[3] + 8), + radius=8, + outline=PALETTE["red"] + (a,), + width=w, + ) + composite = Image.alpha_composite( + composite, glow.filter(ImageFilter.GaussianBlur(4)) + ) draw = ImageDraw.Draw(composite) - draw.rounded_rectangle((bbox[0] - 8, bbox[1] - 7, bbox[2] + 8, bbox[3] + 8), radius=8, outline=PALETTE["red"] + (255,), width=3) + draw.rounded_rectangle( + (bbox[0] - 8, bbox[1] - 7, bbox[2] + 8, bbox[3] + 8), + radius=8, + outline=PALETTE["red"] + (255,), + width=3, + ) - mask_delta = Image.blend(original.convert("RGB"), answer_mask.convert("RGB"), 0.42).convert("RGBA") - crop = mask_delta.crop((max(0, bbox[0] - 76), max(0, bbox[1] - 42), min(side, bbox[2] + 154), min(side, bbox[3] + 48))) + mask_delta = Image.blend( + original.convert("RGB"), answer_mask.convert("RGB"), 0.42 + ).convert("RGBA") + crop = mask_delta.crop( + ( + max(0, bbox[0] - 76), + max(0, bbox[1] - 42), + min(side, bbox[2] + 154), + min(side, bbox[3] + 48), + ) + ) crop = crop.resize((crop.width * 3, crop.height * 3), Image.Resampling.NEAREST) crop_draw = ImageDraw.Draw(crop) scale = 3 @@ -197,10 +252,20 @@ def draw_activation_overlay( cy0 = (bbox[1] - max(0, bbox[1] - 42)) * scale cx1 = (bbox[2] - max(0, bbox[0] - 76)) * scale cy1 = (bbox[3] - max(0, bbox[1] - 42)) * scale - crop_draw.rounded_rectangle((cx0 - 4, cy0 - 4, cx1 + 4, cy1 + 4), radius=8, outline=PALETTE["red"] + (255,), width=5) + crop_draw.rounded_rectangle( + (cx0 - 4, cy0 - 4, cx1 + 4, cy1 + 4), + radius=8, + outline=PALETTE["red"] + (255,), + width=5, + ) composite.alpha_composite(crop, (side - crop.width - 20, 20)) draw = ImageDraw.Draw(composite) - draw.text((side - crop.width - 16, 20 + crop.height + 8), "answer glyph crop: original → masked", fill=PALETTE["ink"] + (235,), font=font(18, True)) + draw.text( + (side - crop.width - 16, 20 + crop.height + 8), + "answer glyph crop: original → masked", + fill=PALETTE["ink"] + (235,), + font=font(18, True), + ) return composite @@ -212,31 +277,72 @@ def draw_layer_bars( answer_region_layer: np.ndarray, ratio_layer: np.ndarray, ) -> None: - rounded_panel(draw, box, "layer-by-layer scar", "red = answer mask, green = equal random mask, cyan = answer glyph tokens") + rounded_panel( + draw, + box, + "layer-by-layer scar", + "red = answer mask, green = equal random mask, cyan = answer glyph tokens", + ) x0, y0, x1, y1 = box chart = (x0 + 74, y0 + 103, x1 - 34, y1 - 72) rows = answer_layer.size row_h = (chart[3] - chart[1]) / rows - scale = float(np.quantile(np.concatenate([answer_layer, random_layer, answer_region_layer]), 0.96)) + scale = float( + np.quantile( + np.concatenate([answer_layer, random_layer, answer_region_layer]), 0.96 + ) + ) scale = max(scale, 1e-6) for i in range(rows): y = chart[1] + i * row_h - draw.text((x0 + 28, round(y + row_h * 0.18)), f"L{i:02d}", fill=PALETTE["muted"], font=font(12)) + draw.text( + (x0 + 28, round(y + row_h * 0.18)), + f"L{i:02d}", + fill=PALETTE["muted"], + font=font(12), + ) max_w = chart[2] - chart[0] aw = round(max_w * min(1.0, float(answer_layer[i]) / scale)) rw = round(max_w * min(1.0, float(random_layer[i]) / scale)) gw = round(max_w * min(1.0, float(answer_region_layer[i]) / scale)) yy = round(y) - draw.rounded_rectangle((chart[0], yy + 2, chart[0] + aw, yy + 8), radius=3, fill=PALETTE["red"]) - draw.rounded_rectangle((chart[0], yy + 11, chart[0] + rw, yy + 17), radius=3, fill=PALETTE["green"]) - draw.rounded_rectangle((chart[0], yy + 20, chart[0] + gw, yy + 27), radius=3, fill=PALETTE["cyan"]) + draw.rounded_rectangle( + (chart[0], yy + 2, chart[0] + aw, yy + 8), radius=3, fill=PALETTE["red"] + ) + draw.rounded_rectangle( + (chart[0], yy + 11, chart[0] + rw, yy + 17), radius=3, fill=PALETTE["green"] + ) + draw.rounded_rectangle( + (chart[0], yy + 20, chart[0] + gw, yy + 27), radius=3, fill=PALETTE["cyan"] + ) ratio = float(ratio_layer[i]) - draw.text((chart[2] - 58, yy + 8), f"{ratio:4.1f}×", fill=PALETTE["amber"], font=font(13, True)) - draw.text((chart[0], y1 - 45), "Mean delta per decoder layer. Ratio labels compare answer-mask vs random-mask deltas.", fill=PALETTE["muted"], font=font(14)) + draw.text( + (chart[2] - 58, yy + 8), + f"{ratio:4.1f}×", + fill=PALETTE["amber"], + font=font(13, True), + ) + draw.text( + (chart[0], y1 - 45), + "Mean delta per decoder layer. Ratio labels compare answer-mask vs random-mask deltas.", + fill=PALETTE["muted"], + font=font(14), + ) -def draw_scar_strip(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], ratio_norm: np.ndarray, answer_tokens: list[int], top_tokens: list[int]) -> None: - rounded_panel(draw, box, "token scar matrix", "decoder layers × image tokens; vertical lines locate answer glyphs and top scar bins") +def draw_scar_strip( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + ratio_norm: np.ndarray, + answer_tokens: list[int], + top_tokens: list[int], +) -> None: + rounded_panel( + draw, + box, + "token scar matrix", + "decoder layers × image tokens; vertical lines locate answer glyphs and top scar bins", + ) x0, y0, x1, y1 = box hx0, hy0, hx1, hy1 = x0 + 58, y0 + 90, x1 - 28, y1 - 54 rows, cols = ratio_norm.shape @@ -258,10 +364,19 @@ def draw_scar_strip(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], r for r in range(0, rows, 4): y = round(hy0 + (r + 0.5) * ch) draw.text((x0 + 22, y - 7), str(r), fill=PALETTE["muted"], font=font(12)) - draw.text((hx0, y1 - 32), "image-token sequence →", fill=PALETTE["muted"], font=font(13)) + draw.text( + (hx0, y1 - 32), "image-token sequence →", fill=PALETTE["muted"], font=font(13) + ) -def draw_top_token_table(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], top_tokens: list[int], ratio_mean: np.ndarray, answer_mean: np.ndarray, grid: int) -> None: +def draw_top_token_table( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + top_tokens: list[int], + ratio_mean: np.ndarray, + answer_mean: np.ndarray, + grid: int, +) -> None: rounded_panel(draw, box, "highest-scar token bins", "actual heatmaps.npz token IDs") x0, y0, _, y1 = box y = y0 + 90 @@ -270,12 +385,28 @@ def draw_top_token_table(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, in bar_scale = max(1e-6, float(np.quantile(answer_mean, 0.98))) for rank, tok in enumerate(top_tokens[:max_rows], start=1): r, c = divmod(tok, grid) - draw.text((x0 + 28, y), f"{rank:02d}", fill=PALETTE["amber"], font=font(13, True)) - draw.text((x0 + 68, y), f"token {tok:03d}", fill=PALETTE["ink"], font=font(14, True)) - draw.text((x0 + 164, y), f"grid r{r:02d} c{c:02d}", fill=PALETTE["muted"], font=font(13)) - draw.text((x0 + 292, y), f"ratio {ratio_mean[tok]:.2f}×", fill=PALETTE["cyan"], font=font(13, True)) + draw.text( + (x0 + 28, y), f"{rank:02d}", fill=PALETTE["amber"], font=font(13, True) + ) + draw.text( + (x0 + 68, y), f"token {tok:03d}", fill=PALETTE["ink"], font=font(14, True) + ) + draw.text( + (x0 + 164, y), + f"grid r{r:02d} c{c:02d}", + fill=PALETTE["muted"], + font=font(13), + ) + draw.text( + (x0 + 292, y), + f"ratio {ratio_mean[tok]:.2f}×", + fill=PALETTE["cyan"], + font=font(13, True), + ) bar_w = round(112 * min(1.0, float(answer_mean[tok]) / bar_scale)) - draw.rounded_rectangle((x0 + 408, y + 4, x0 + 408 + bar_w, y + 14), radius=4, fill=PALETTE["red"]) + draw.rounded_rectangle( + (x0 + 408, y + 4, x0 + 408 + bar_w, y + 14), radius=4, fill=PALETTE["red"] + ) y += row_gap @@ -303,24 +434,39 @@ def render(source: Path, out_dir: Path) -> None: grid = int(round(math.sqrt(token_count))) boxes = token_boxes(original.width, grid) answer_area = (bbox[0], bbox[1], bbox[2], bbox[3]) - answer_tokens = [i for i, b in enumerate(boxes) if intersect_area(answer_area, b) > 0] + answer_tokens = [ + i for i, b in enumerate(boxes) if intersect_area(answer_area, b) > 0 + ] if not answer_tokens: center_x = (bbox[0] + bbox[2]) / 2 center_y = (bbox[1] + bbox[3]) / 2 - answer_tokens = [min(token_count - 1, max(0, int(center_y / original.height * grid) * grid + int(center_x / original.width * grid)))] + answer_tokens = [ + min( + token_count - 1, + max( + 0, + int(center_y / original.height * grid) * grid + + int(center_x / original.width * grid), + ), + ) + ] ratio_mean = ratio.mean(axis=0) answer_mean = answer_delta.mean(axis=0) token_score = quantile_norm(ratio_mean, 0.985) answer_set = set(answer_tokens) - top_tokens = [int(i) for i in np.argsort(ratio_mean)[::-1] if int(i) not in answer_set][:24] + top_tokens = [ + int(i) for i in np.argsort(ratio_mean)[::-1] if int(i) not in answer_set + ][:24] answer_region_layer = answer_delta[:, answer_tokens].mean(axis=1) answer_layer = answer_delta.mean(axis=1) random_layer = random_delta.mean(axis=1) ratio_layer = answer_layer / np.maximum(random_layer, 1e-6) out_dir.mkdir(parents=True, exist_ok=True) - overlay = draw_activation_overlay(original, answer_mask, token_score, top_tokens[:18], answer_tokens, bbox) + overlay = draw_activation_overlay( + original, answer_mask, token_score, top_tokens[:18], answer_tokens, bbox + ) overlay = overlay.resize((760, 760), Image.Resampling.LANCZOS) W, H = 1900, 1260 @@ -336,20 +482,55 @@ def render(source: Path, out_dir: Path) -> None: canvas = Image.alpha_composite(canvas, glow.filter(ImageFilter.GaussianBlur(78))) draw = ImageDraw.Draw(canvas) - draw.text((64, 42), "SNAPCOMPACT GLYPH MATRIX", fill=PALETTE["amber"], font=font(22, True)) - draw.text((64, 78), "The answer glyphs leave a hidden activation scar", fill=PALETTE["ink"], font=font(56, True)) + draw.text( + (64, 42), "SNAPCOMPACT GLYPH MATRIX", fill=PALETTE["amber"], font=font(22, True) + ) + draw.text( + (64, 78), + "The answer glyphs leave a hidden activation scar", + fill=PALETTE["ink"], + font=font(56, True), + ) subtitle = "Original dense text bitmap, overlaid with answer/random activation ratios from 19 decoder layers × 729 image tokens." draw.text((68, 145), subtitle, fill=PALETTE["muted"], font=font(22)) - rounded_panel(draw, (52, 205, 862, 1066), "visible glyphs ↔ hidden tokens", "red box = actual answer cells; cyan = intersecting image tokens; amber = top scar bins") + rounded_panel( + draw, + (52, 205, 862, 1066), + "visible glyphs ↔ hidden tokens", + "red box = actual answer cells; cyan = intersecting image tokens; amber = top scar bins", + ) paste_shadowed(canvas, overlay, (78, 282)) - draw.text((82, 1085), f"Question: {q['q']}", fill=PALETTE["ink"], font=font(21, True)) - draw.text((82, 1120), f"Gold answer: {q['answer_text']} · cells {q['answer_start']}–{q['answer_end'] - 1}", fill=PALETTE["amber"], font=font(24, True)) - draw.text((82, 1160), f"Answer/random mean delta: {summary['answer_over_random_delta']:.2f}×", fill=PALETTE["cyan"], font=font(22, True)) + draw.text( + (82, 1085), f"Question: {q['q']}", fill=PALETTE["ink"], font=font(21, True) + ) + draw.text( + (82, 1120), + f"Gold answer: {q['answer_text']} · cells {q['answer_start']}–{q['answer_end'] - 1}", + fill=PALETTE["amber"], + font=font(24, True), + ) + draw.text( + (82, 1160), + f"Answer/random mean delta: {summary['answer_over_random_delta']:.2f}×", + fill=PALETTE["cyan"], + font=font(22, True), + ) - draw_layer_bars(draw, (900, 205, 1838, 628), answer_layer, random_layer, answer_region_layer, ratio_layer) - draw_scar_strip(draw, (900, 662, 1838, 930), ratio_norm_binned, answer_tokens, top_tokens) - draw_top_token_table(draw, (900, 964, 1838, 1196), top_tokens, ratio_mean, answer_mean, grid) + draw_layer_bars( + draw, + (900, 205, 1838, 628), + answer_layer, + random_layer, + answer_region_layer, + ratio_layer, + ) + draw_scar_strip( + draw, (900, 662, 1838, 930), ratio_norm_binned, answer_tokens, top_tokens + ) + draw_top_token_table( + draw, (900, 964, 1838, 1196), top_tokens, ratio_mean, answer_mean, grid + ) for i in range(240): draw.rectangle((1568 + i, 156, 1569 + i, 174), fill=heat_color(i / 239)) @@ -363,7 +544,13 @@ def render(source: Path, out_dir: Path) -> None: source_data = { "source": str(source), "question": q, - "geometry": {"text_cols": cols, "text_rows": rows, "glyph_adv": adv, "glyph_pitch": pitch, "image_token_grid": [grid, grid]}, + "geometry": { + "text_cols": cols, + "text_rows": rows, + "glyph_adv": adv, + "glyph_pitch": pitch, + "image_token_grid": [grid, grid], + }, "answer_bbox_pixels": list(map(int, bbox)), "answer_image_tokens": [int(x) for x in answer_tokens], "top_scar_tokens": [ diff --git a/packages/snapcompact/research/snapcompact_viz_radial.py b/packages/snapcompact/research/snapcompact_viz_radial.py index a5f477d8c..926a62017 100644 --- a/packages/snapcompact/research/snapcompact_viz_radial.py +++ b/packages/snapcompact/research/snapcompact_viz_radial.py @@ -59,7 +59,9 @@ def radar_cmap() -> mcolors.LinearSegmentedColormap: return mcolors.LinearSegmentedColormap.from_list("snapcompact_radar", colors) -def top_echoes(ratio: np.ndarray, answer: np.ndarray, random: np.ndarray, limit: int = 18) -> list[dict[str, float | int]]: +def top_echoes( + ratio: np.ndarray, answer: np.ndarray, random: np.ndarray, limit: int = 18 +) -> list[dict[str, float | int]]: flat = np.argpartition(ratio.ravel(), -limit)[-limit:] flat = flat[np.argsort(ratio.ravel()[flat])[::-1]] rows: list[dict[str, float | int]] = [] @@ -70,7 +72,9 @@ def top_echoes(ratio: np.ndarray, answer: np.ndarray, random: np.ndarray, limit: "rank": len(rows) + 1, "layer": int(layer), "bin": int(bin_idx), - "angle_degrees": round(float((bin_idx + 0.5) * 360.0 / ratio.shape[1]), 2), + "angle_degrees": round( + float((bin_idx + 0.5) * 360.0 / ratio.shape[1]), 2 + ), "answer_delta": round(float(answer[layer, bin_idx]), 4), "random_delta": round(float(random[layer, bin_idx]), 4), "answer_random_ratio": round(float(ratio[layer, bin_idx]), 4), @@ -93,11 +97,26 @@ def add_glow_spikes(ax: plt.Axes, ratio: np.ndarray, norm_ratio: np.ndarray) -> tip_r = base_r + 0.12 + 0.58 * v theta = float(theta_centers[idx]) color = AMBER if v > 0.78 else CYAN - ax.plot([theta, theta], [base_r, tip_r], color=color, linewidth=0.7 + 1.8 * v, alpha=0.30 + 0.55 * v) - ax.scatter([theta], [tip_r], s=5 + 28 * v, color=color, alpha=0.26 + 0.55 * v, linewidths=0) + ax.plot( + [theta, theta], + [base_r, tip_r], + color=color, + linewidth=0.7 + 1.8 * v, + alpha=0.30 + 0.55 * v, + ) + ax.scatter( + [theta], + [tip_r], + s=5 + 28 * v, + color=color, + alpha=0.26 + 0.55 * v, + linewidths=0, + ) -def draw_radial(summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np.ndarray) -> plt.Figure: +def draw_radial( + summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np.ndarray +) -> plt.Figure: layers, bins = ratio.shape norm_ratio = robust_norm(ratio, 0.972) theta_edges, radius_edges = polar_edges(bins, layers) @@ -116,7 +135,14 @@ def draw_radial(summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np ax.spines["polar"].set_color("#38f4ff66") ax.spines["polar"].set_linewidth(1.2) - ax.pcolormesh(theta_grid, radius_grid, norm_ratio, cmap=radar_cmap(), shading="flat", alpha=0.96) + ax.pcolormesh( + theta_grid, + radius_grid, + norm_ratio, + cmap=radar_cmap(), + shading="flat", + alpha=0.96, + ) # Soft trace underneath the hottest angular bearings, like phosphor persistence. bearing_strength = norm_ratio.mean(axis=0) + norm_ratio.max(axis=0) * 0.42 @@ -126,18 +152,43 @@ def draw_radial(summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np for width, alpha in ((38, 0.055), (22, 0.075), (8, 0.14)): half_width = np.deg2rad(width / 2) theta = np.linspace(sweep_theta - half_width, sweep_theta + half_width, 80) - ax.fill_between(theta, 0.0, layers + 1.75, color=GREEN, alpha=alpha, linewidth=0) + ax.fill_between( + theta, 0.0, layers + 1.75, color=GREEN, alpha=alpha, linewidth=0 + ) add_glow_spikes(ax, ratio, norm_ratio) for r in range(1, layers + 2): - ax.plot(np.linspace(0, 2 * np.pi, 360), np.full(360, r), color="#6fffe522", linewidth=0.55) + ax.plot( + np.linspace(0, 2 * np.pi, 360), + np.full(360, r), + color="#6fffe522", + linewidth=0.55, + ) for deg in range(0, 360, 15): th = np.deg2rad(deg) ax.plot([th, th], [1, layers + 1.4], color="#6fffe516", linewidth=0.45) - ax.text(0.5, 0.5, "ECHO\nCORE", color="#dff", fontsize=13, fontweight="bold", ha="center", va="center", transform=ax.transAxes) - ax.text(np.deg2rad(sweep_angle), layers + 1.35, "strongest bearing", color=GREEN, fontsize=8, ha="center", va="center") + ax.text( + 0.5, + 0.5, + "ECHO\nCORE", + color="#dff", + fontsize=13, + fontweight="bold", + ha="center", + va="center", + transform=ax.transAxes, + ) + ax.text( + np.deg2rad(sweep_angle), + layers + 1.35, + "strongest bearing", + color=GREEN, + fontsize=8, + ha="center", + va="center", + ) side = fig.add_axes([0.70, 0.06, 0.27, 0.86], facecolor=BG) side.axis("off") @@ -150,9 +201,35 @@ def draw_radial(summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np max_echo = top[0] title_fx = [pe.withStroke(linewidth=4, foreground="#0b1918")] - side.text(0.00, 0.98, "SNAPCOMPACT RADAR", color=GREEN, fontsize=12, fontweight="bold", va="top") - side.text(0.00, 0.925, "Where the missing\nanswer echoes", color=INK, fontsize=27, fontweight="bold", va="top", linespacing=0.92, path_effects=title_fx) - side.text(0.00, 0.765, "Concentric rings are decoder layers. Angles are image-token bins. Bright spikes are answer-mask residuals divided by the random-mask control.", color=MUTED, fontsize=9.5, va="top", wrap=True) + side.text( + 0.00, + 0.98, + "SNAPCOMPACT RADAR", + color=GREEN, + fontsize=12, + fontweight="bold", + va="top", + ) + side.text( + 0.00, + 0.925, + "Where the missing\nanswer echoes", + color=INK, + fontsize=27, + fontweight="bold", + va="top", + linespacing=0.92, + path_effects=title_fx, + ) + side.text( + 0.00, + 0.765, + "Concentric rings are decoder layers. Angles are image-token bins. Bright spikes are answer-mask residuals divided by the random-mask control.", + color=MUTED, + fontsize=9.5, + va="top", + wrap=True, + ) metrics = [ ("gold answer", str(q["answer_text"]), AMBER), @@ -161,21 +238,68 @@ def draw_radial(summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np ("layers", f"{summary['layers']}", CYAN), ("mean answer/random Δ", f"{ratio_mean:.2f}×", AMBER), ("max-ratio layer", f"L{max_layer}", GREEN), - ("loudest echo", f"L{max_echo['layer']} · bin {max_echo['bin']} · {max_echo['answer_random_ratio']:.1f}×", RED), + ( + "loudest echo", + f"L{max_echo['layer']} · bin {max_echo['bin']} · {max_echo['answer_random_ratio']:.1f}×", + RED, + ), ] y = 0.655 for label, value, color in metrics: - side.text(0.00, y, label.upper(), color=MUTED, fontsize=7.2, fontweight="bold", va="top") + side.text( + 0.00, + y, + label.upper(), + color=MUTED, + fontsize=7.2, + fontweight="bold", + va="top", + ) value_size = 12.6 if len(value) < 34 else 8.7 - side.text(0.00, y - 0.026, value, color=color, fontsize=value_size, fontweight="bold" if label != "question" else "normal", va="top", wrap=True) + side.text( + 0.00, + y - 0.026, + value, + color=color, + fontsize=value_size, + fontweight="bold" if label != "question" else "normal", + va="top", + wrap=True, + ) y -= 0.075 if label != "question" else 0.105 - side.text(0.00, y - 0.006, "TOP ECHOES", color=GREEN, fontsize=7.6, fontweight="bold", va="top") + side.text( + 0.00, + y - 0.006, + "TOP ECHOES", + color=GREEN, + fontsize=7.6, + fontweight="bold", + va="top", + ) y -= 0.040 for row in top[:4]: - intensity = min(1.0, float(row["answer_random_ratio"]) / float(max_echo["answer_random_ratio"])) - side.plot([0.00, 0.36 * intensity], [y - 0.004, y - 0.004], color=AMBER, linewidth=3.2, alpha=0.35 + 0.55 * intensity, solid_capstyle="round") - side.text(0.40, y - 0.014, f"L{row['layer']:02d} bin {row['bin']:03d} {row['answer_random_ratio']:>5.1f}×", color=INK, fontsize=7.4, va="bottom", family="monospace") + intensity = min( + 1.0, + float(row["answer_random_ratio"]) / float(max_echo["answer_random_ratio"]), + ) + side.plot( + [0.00, 0.36 * intensity], + [y - 0.004, y - 0.004], + color=AMBER, + linewidth=3.2, + alpha=0.35 + 0.55 * intensity, + solid_capstyle="round", + ) + side.text( + 0.40, + y - 0.014, + f"L{row['layer']:02d} bin {row['bin']:03d} {row['answer_random_ratio']:>5.1f}×", + color=INK, + fontsize=7.4, + va="bottom", + family="monospace", + ) y -= 0.032 # Tiny color scale and data provenance line. @@ -184,7 +308,13 @@ def draw_radial(summary: dict, answer: np.ndarray, random: np.ndarray, ratio: np grad_ax.set_axis_off() side.text(0.00, 0.006, "low ratio", color=MUTED, fontsize=7, va="bottom") side.text(0.59, 0.006, "high answer echo", color=MUTED, fontsize=7, va="bottom") - fig.text(0.045, 0.018, "Actual heatmaps.npz arrays: ratio_binned, answer_binned, random_binned", color="#8da0a888", fontsize=8) + fig.text( + 0.045, + 0.018, + "Actual heatmaps.npz arrays: ratio_binned, answer_binned, random_binned", + color="#8da0a888", + fontsize=8, + ) return fig @@ -210,8 +340,19 @@ def main() -> None: plt.close(fig) echoes = top_echoes(ratio, answer, random, 24) - (out_dir / "radial_top_echoes.json").write_text(json.dumps({"source": str(data_dir / "heatmaps.npz"), "top_echoes": echoes}, indent=2) + "\n") - np.savez_compressed(out_dir / "radial_source.npz", answer_binned=answer, random_binned=random, ratio_binned=ratio, ratio_norm=robust_norm(ratio, 0.972)) + (out_dir / "radial_top_echoes.json").write_text( + json.dumps( + {"source": str(data_dir / "heatmaps.npz"), "top_echoes": echoes}, indent=2 + ) + + "\n" + ) + np.savez_compressed( + out_dir / "radial_source.npz", + answer_binned=answer, + random_binned=random, + ratio_binned=ratio, + ratio_norm=robust_norm(ratio, 0.972), + ) print(out_png) diff --git a/packages/snapcompact/research/snapcompact_viz_token_grid.py b/packages/snapcompact/research/snapcompact_viz_token_grid.py index 7087390ad..a6b476220 100644 --- a/packages/snapcompact/research/snapcompact_viz_token_grid.py +++ b/packages/snapcompact/research/snapcompact_viz_token_grid.py @@ -41,9 +41,15 @@ PALETTE = { def font(size: int, bold: bool = False) -> ImageFont.FreeTypeFont | ImageFont.ImageFont: candidates = [ - "/System/Library/Fonts/Supplemental/Arial Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Arial.ttf", - "/System/Library/Fonts/Supplemental/Helvetica Bold.ttf" if bold else "/System/Library/Fonts/Supplemental/Helvetica.ttf", - "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" if bold else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/System/Library/Fonts/Supplemental/Arial Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Arial.ttf", + "/System/Library/Fonts/Supplemental/Helvetica Bold.ttf" + if bold + else "/System/Library/Fonts/Supplemental/Helvetica.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" + if bold + else "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", ] for candidate in candidates: if candidate and Path(candidate).exists(): @@ -55,7 +61,9 @@ def lerp(a: int, b: int, t: float) -> int: return round(a + (b - a) * t) -def mix(a: tuple[int, int, int], b: tuple[int, int, int], t: float) -> tuple[int, int, int]: +def mix( + a: tuple[int, int, int], b: tuple[int, int, int], t: float +) -> tuple[int, int, int]: return (lerp(a[0], b[0], t), lerp(a[1], b[1], t), lerp(a[2], b[2], t)) @@ -89,7 +97,12 @@ def token_side(summary: dict, token_count: int) -> int: grid = summary.get("processor_meta", {}).get("image_grid_thw", [[1, 0, 0]])[0] _, gh, gw = grid merge = math.isqrt(max(1, (gh * gw) // token_count)) - if merge and gh % merge == 0 and gw % merge == 0 and (gh // merge) * (gw // merge) == token_count: + if ( + merge + and gh % merge == 0 + and gw % merge == 0 + and (gh // merge) * (gw // merge) == token_count + ): return gh // merge raise ValueError(f"cannot fold {token_count} image tokens into a square grid") @@ -100,7 +113,9 @@ def fold_tokens(arr: np.ndarray, side: int) -> np.ndarray: return arr.reshape(arr.shape[0], side, side) -def heat_overlay(base: Image.Image, heat: np.ndarray, alpha_floor: int = 28, alpha_peak: int = 220) -> Image.Image: +def heat_overlay( + base: Image.Image, heat: np.ndarray, alpha_floor: int = 28, alpha_peak: int = 220 +) -> Image.Image: norm, _ = normalize(heat) small = Image.new("RGBA", (heat.shape[1], heat.shape[0]), (0, 0, 0, 0)) pix = small.load() @@ -108,13 +123,27 @@ def heat_overlay(base: Image.Image, heat: np.ndarray, alpha_floor: int = 28, alp for x in range(heat.shape[1]): t = float(norm[y, x]) r, g, b = heat_color(t) - pix[x, y] = (r, g, b, round(alpha_floor + (alpha_peak - alpha_floor) * (t ** 0.85))) - overlay = small.resize(base.size, Image.Resampling.BICUBIC).filter(ImageFilter.GaussianBlur(1.0)) - dim = Image.blend(base.convert("RGB"), Image.new("RGB", base.size, (5, 8, 13)), 0.28).convert("RGBA") + pix[x, y] = ( + r, + g, + b, + round(alpha_floor + (alpha_peak - alpha_floor) * (t**0.85)), + ) + overlay = small.resize(base.size, Image.Resampling.BICUBIC).filter( + ImageFilter.GaussianBlur(1.0) + ) + dim = Image.blend( + base.convert("RGB"), Image.new("RGB", base.size, (5, 8, 13)), 0.28 + ).convert("RGBA") return Image.alpha_composite(dim, overlay).convert("RGB") -def draw_token_grid(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], side: int, color: tuple[int, int, int] = (255, 255, 255)) -> None: +def draw_token_grid( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + side: int, + color: tuple[int, int, int] = (255, 255, 255), +) -> None: x0, y0, x1, y1 = box for i in range(side + 1): x = round(x0 + (x1 - x0) * i / side) @@ -124,7 +153,12 @@ def draw_token_grid(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], s draw.line((x0, y, x1, y), fill=fill, width=1) -def paste_fit(canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, int], resample: int = Image.Resampling.LANCZOS) -> tuple[int, int, int, int]: +def paste_fit( + canvas: Image.Image, + img: Image.Image, + box: tuple[int, int, int, int], + resample: int = Image.Resampling.LANCZOS, +) -> tuple[int, int, int, int]: x0, y0, x1, y1 = box scale = min((x1 - x0) / img.width, (y1 - y0) / img.height) w = max(1, round(img.width * scale)) @@ -136,7 +170,15 @@ def paste_fit(canvas: Image.Image, img: Image.Image, box: tuple[int, int, int, i return (px, py, px + w, py + h) -def crop_answer(img: Image.Image, start: int, end: int, cols: int, adv: int, pitch: int, pad_cells: int = 34) -> Image.Image: +def crop_answer( + img: Image.Image, + start: int, + end: int, + cols: int, + adv: int, + pitch: int, + pad_cells: int = 34, +) -> Image.Image: rows = img.height // pitch row0 = max(0, start // cols - 5) row1 = min(rows, end // cols + 6) @@ -154,7 +196,9 @@ def crop_answer(img: Image.Image, start: int, end: int, cols: int, adv: int, pit return crop -def answer_bbox(start: int, end: int, cols: int, adv: int, pitch: int) -> tuple[int, int, int, int]: +def answer_bbox( + start: int, end: int, cols: int, adv: int, pitch: int +) -> tuple[int, int, int, int]: return ( max(0, (start % cols) * adv - adv), max(0, (start // cols) * pitch - 2), @@ -163,15 +207,28 @@ def answer_bbox(start: int, end: int, cols: int, adv: int, pitch: int) -> tuple[ ) -def draw_panel(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], title: str, subtitle: str | None = None) -> None: - draw.rounded_rectangle(box, radius=26, fill=PALETTE["panel"], outline=(32, 43, 55), width=1) +def draw_panel( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + title: str, + subtitle: str | None = None, +) -> None: + draw.rounded_rectangle( + box, radius=26, fill=PALETTE["panel"], outline=(32, 43, 55), width=1 + ) x0, y0, _, _ = box draw.text((x0 + 24, y0 + 20), title, fill=PALETTE["ink"], font=font(28, True)) if subtitle: draw.text((x0 + 24, y0 + 56), subtitle, fill=PALETTE["muted"], font=font(17)) -def draw_micro_grid(canvas: Image.Image, heat: np.ndarray, box: tuple[int, int, int, int], title: str, subtitle: str) -> None: +def draw_micro_grid( + canvas: Image.Image, + heat: np.ndarray, + box: tuple[int, int, int, int], + title: str, + subtitle: str, +) -> None: draw = ImageDraw.Draw(canvas) draw_panel(draw, box, title, subtitle) x0, y0, x1, y1 = box @@ -194,16 +251,34 @@ def draw_micro_grid(canvas: Image.Image, heat: np.ndarray, box: tuple[int, int, draw.line((gx0, y, gx1, y), fill=(255, 255, 255, 34)) -def label(draw: ImageDraw.ImageDraw, xy: tuple[int, int], text: str, color: tuple[int, int, int], size: int = 18, bold: bool = True) -> None: +def label( + draw: ImageDraw.ImageDraw, + xy: tuple[int, int], + text: str, + color: tuple[int, int, int], + size: int = 18, + bold: bool = True, +) -> None: x, y = xy pad = 8 f = font(size, bold) box = draw.textbbox((x, y), text, font=f) - draw.rounded_rectangle((box[0] - pad, box[1] - 4, box[2] + pad, box[3] + 5), radius=9, fill=(4, 6, 10), outline=color, width=1) + draw.rounded_rectangle( + (box[0] - pad, box[1] - 4, box[2] + pad, box[3] + 5), + radius=9, + fill=(4, 6, 10), + outline=color, + width=1, + ) draw.text((x, y), text, fill=color, font=f) -def draw_hotspots(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], heat: np.ndarray, count: int = 9) -> None: +def draw_hotspots( + draw: ImageDraw.ImageDraw, + box: tuple[int, int, int, int], + heat: np.ndarray, + count: int = 9, +) -> None: x0, y0, x1, y1 = box side = heat.shape[0] flat = heat.ravel() @@ -211,7 +286,10 @@ def draw_hotspots(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], hea chosen: list[int] = [] for idx in np.argsort(flat)[::-1]: r, c = divmod(int(idx), side) - if all(abs(r - divmod(j, side)[0]) + abs(c - divmod(j, side)[1]) >= 3 for j in chosen): + if all( + abs(r - divmod(j, side)[0]) + abs(c - divmod(j, side)[1]) >= 3 + for j in chosen + ): chosen.append(int(idx)) if len(chosen) == count: break @@ -220,12 +298,26 @@ def draw_hotspots(draw: ImageDraw.ImageDraw, box: tuple[int, int, int, int], hea cx = round(x0 + (c + 0.5) * (x1 - x0) / side) cy = round(y0 + (r + 0.5) * (y1 - y0) / side) rad = 11 if rank <= 3 else 8 - draw.ellipse((cx - rad, cy - rad, cx + rad, cy + rad), outline=PALETTE["amber"], width=3) + draw.ellipse( + (cx - rad, cy - rad, cx + rad, cy + rad), outline=PALETTE["amber"], width=3 + ) if rank <= 5: - draw.text((cx + 10, cy - 16), str(rank), fill=PALETTE["amber"], font=font(16, True)) + draw.text( + (cx + 10, cy - 16), + str(rank), + fill=PALETTE["amber"], + font=font(16, True), + ) -def text_block(draw: ImageDraw.ImageDraw, xy: tuple[int, int], lines: Iterable[str], fill: tuple[int, int, int], size: int = 20, gap: int = 8) -> None: +def text_block( + draw: ImageDraw.ImageDraw, + xy: tuple[int, int], + lines: Iterable[str], + fill: tuple[int, int, int], + size: int = 20, + gap: int = 8, +) -> None: x, y = xy f = font(size) for line in lines: @@ -263,7 +355,9 @@ def render() -> None: early_ratio=early_ratio, mid_answer_delta=mid_delta, late_answer_delta=late_delta, - image_grid_thw=np.array(summary["processor_meta"]["image_grid_thw"][0], dtype=np.int32), + image_grid_thw=np.array( + summary["processor_meta"]["image_grid_thw"][0], dtype=np.int32 + ), ) (OUT_DIR / "token_grid_summary.json").write_text( json.dumps( @@ -272,7 +366,9 @@ def render() -> None: "image_grid_thw": summary["processor_meta"]["image_grid_thw"][0], "image_tokens": int(summary["image_tokens"]), "rendered_token_grid": [side, side], - "patch_merge": int(summary["processor_meta"]["image_grid_thw"][0][1] // side), + "patch_merge": int( + summary["processor_meta"]["image_grid_thw"][0][1] // side + ), "answer_over_random_delta": float(summary["answer_over_random_delta"]), "question": summary["question"]["q"], "answer_text": summary["question"]["answer_text"], @@ -291,11 +387,20 @@ def render() -> None: gd.ellipse((-320, -240, 960, 780), fill=(255, 80, 66, 34)) gd.ellipse((920, -120, 2350, 1100), fill=(75, 218, 255, 30)) gd.ellipse((760, 860, 1810, 1760), fill=(255, 194, 72, 18)) - canvas = Image.alpha_composite(canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(95))).convert("RGB") + canvas = Image.alpha_composite( + canvas.convert("RGBA"), glow.filter(ImageFilter.GaussianBlur(95)) + ).convert("RGB") draw = ImageDraw.Draw(canvas) - draw.text((64, 48), "SNAPCOMPACT TOKEN FIELD", fill=PALETTE["amber"], font=font(24, True)) - draw.text((64, 86), "Where the hidden-state scar lands on the bitmap", fill=PALETTE["ink"], font=font(62, True)) + draw.text( + (64, 48), "SNAPCOMPACT TOKEN FIELD", fill=PALETTE["amber"], font=font(24, True) + ) + draw.text( + (64, 86), + "Where the hidden-state scar lands on the bitmap", + fill=PALETTE["ink"], + font=font(62, True), + ) draw.text( (66, 164), "PaddleOCR-VL reports a 1×54×54 image patch grid; 729 hidden-state image tokens fold back to 27×27 spatial cells.", @@ -305,7 +410,12 @@ def render() -> None: # Main spatial map. main_panel = (545, 225, 1455, 1340) - draw_panel(draw, main_panel, "answer-mask delta projected onto image tokens", "mean ||hidden(original) − hidden(answer-mask)|| across 19 layers") + draw_panel( + draw, + main_panel, + "answer-mask delta projected onto image tokens", + "mean ||hidden(original) − hidden(answer-mask)|| across 19 layers", + ) map_box = (610, 330, 1390, 1110) projected = heat_overlay(original, answer_mean) pasted = paste_fit(canvas, projected, map_box, Image.Resampling.LANCZOS) @@ -313,7 +423,13 @@ def render() -> None: overlay = Image.new("RGBA", canvas.size, (0, 0, 0, 0)) od = ImageDraw.Draw(overlay) draw_token_grid(od, pasted, side, (255, 255, 255)) - bbox = answer_bbox(summary["question"]["answer_start"], summary["question"]["answer_end"], summary["geometry"]["cols"], 8, 13) + bbox = answer_bbox( + summary["question"]["answer_start"], + summary["question"]["answer_end"], + summary["geometry"]["cols"], + 8, + 13, + ) sx = (pasted[2] - pasted[0]) / original.width sy = (pasted[3] - pasted[1]) / original.height answer_rect = ( @@ -326,8 +442,20 @@ def render() -> None: draw_hotspots(od, pasted, answer_mean) canvas = Image.alpha_composite(canvas.convert("RGBA"), overlay).convert("RGB") draw = ImageDraw.Draw(canvas) - label(draw, (pasted[0] + 18, pasted[1] + 18), "27×27 reconstructed image-token grid", PALETTE["cyan"], 19) - label(draw, (answer_rect[2] + 14, answer_rect[1] - 5), "erased answer text", PALETTE["red"], 18) + label( + draw, + (pasted[0] + 18, pasted[1] + 18), + "27×27 reconstructed image-token grid", + PALETTE["cyan"], + 19, + ) + label( + draw, + (answer_rect[2] + 14, answer_rect[1] - 5), + "erased answer text", + PALETTE["red"], + 18, + ) text_block( draw, (620, 1162), @@ -344,13 +472,39 @@ def render() -> None: # Evidence crops. left = (64, 225, 505, 1340) draw_panel(draw, left, "bitmap intervention", "original crop vs. answer erased") - crop = crop_answer(original, summary["question"]["answer_start"], summary["question"]["answer_end"], summary["geometry"]["cols"], 8, 13) - mcrop = crop_answer(masked, summary["question"]["answer_start"], summary["question"]["answer_end"], summary["geometry"]["cols"], 8, 13) + crop = crop_answer( + original, + summary["question"]["answer_start"], + summary["question"]["answer_end"], + summary["geometry"]["cols"], + 8, + 13, + ) + mcrop = crop_answer( + masked, + summary["question"]["answer_start"], + summary["question"]["answer_end"], + summary["geometry"]["cols"], + 8, + 13, + ) draw.text((96, 332), "ORIGINAL", fill=PALETTE["cyan"], font=font(17, True)) - draw.rounded_rectangle((94, 360, 475, 525), radius=16, fill=(240, 238, 226), outline=PALETTE["cyan"], width=3) + draw.rounded_rectangle( + (94, 360, 475, 525), + radius=16, + fill=(240, 238, 226), + outline=PALETTE["cyan"], + width=3, + ) paste_fit(canvas, crop, (108, 374, 461, 511), Image.Resampling.NEAREST) draw.text((96, 572), "ANSWER MASK", fill=PALETTE["red"], font=font(17, True)) - draw.rounded_rectangle((94, 600, 475, 765), radius=16, fill=(240, 238, 226), outline=PALETTE["red"], width=3) + draw.rounded_rectangle( + (94, 600, 475, 765), + radius=16, + fill=(240, 238, 226), + outline=PALETTE["red"], + width=3, + ) paste_fit(canvas, mcrop, (108, 614, 461, 751), Image.Resampling.NEAREST) draw.text((96, 822), "source arrays", fill=PALETTE["muted"], font=font(17, True)) text_block( @@ -370,17 +524,59 @@ def render() -> None: 21, 8, ) - draw.rounded_rectangle((96, 1110, 472, 1268), radius=18, fill=PALETTE["panel2"], outline=(38, 51, 64), width=1) + draw.rounded_rectangle( + (96, 1110, 472, 1268), + radius=18, + fill=PALETTE["panel2"], + outline=(38, 51, 64), + width=1, + ) draw.text((118, 1132), "scar strength", fill=PALETTE["amber"], font=font(18, True)) - draw.text((118, 1170), f"{summary['answer_over_random_delta']:.2f}×", fill=PALETTE["ink"], font=font(54, True)) - draw.text((120, 1232), "answer-mask / random-mask mean delta", fill=PALETTE["muted"], font=font(17)) + draw.text( + (118, 1170), + f"{summary['answer_over_random_delta']:.2f}×", + fill=PALETTE["ink"], + font=font(54, True), + ) + draw.text( + (120, 1232), + "answer-mask / random-mask mean delta", + fill=PALETTE["muted"], + font=font(17), + ) # Right analytical small multiples. - draw_micro_grid(canvas, ratio_mean, (1495, 225, 2136, 590), "ratio field", "mean answer_delta / random_delta") - draw_micro_grid(canvas, early_ratio, (1495, 620, 1810, 975), "early layers", "ratio, layers 0–3") - draw_micro_grid(canvas, mid_delta, (1820, 620, 2136, 975), "middle layers", "answer delta, layers 6–12") - draw_micro_grid(canvas, late_delta, (1495, 1005, 1810, 1340), "late layers", "answer delta, last 4") - draw_micro_grid(canvas, random_mean, (1820, 1005, 2136, 1340), "random control", "random-mask delta") + draw_micro_grid( + canvas, + ratio_mean, + (1495, 225, 2136, 590), + "ratio field", + "mean answer_delta / random_delta", + ) + draw_micro_grid( + canvas, early_ratio, (1495, 620, 1810, 975), "early layers", "ratio, layers 0–3" + ) + draw_micro_grid( + canvas, + mid_delta, + (1820, 620, 2136, 975), + "middle layers", + "answer delta, layers 6–12", + ) + draw_micro_grid( + canvas, + late_delta, + (1495, 1005, 1810, 1340), + "late layers", + "answer delta, last 4", + ) + draw_micro_grid( + canvas, + random_mean, + (1820, 1005, 2136, 1340), + "random control", + "random-mask delta", + ) # Color legend. lx0, ly0, lx1, ly1 = 1530, 530, 2100, 552 diff --git a/packages/snapcompact/research/snapcompact_viz_volume.py b/packages/snapcompact/research/snapcompact_viz_volume.py index 7089b3608..b11d52d01 100644 --- a/packages/snapcompact/research/snapcompact_viz_volume.py +++ b/packages/snapcompact/research/snapcompact_viz_volume.py @@ -44,7 +44,9 @@ def robust01(values: np.ndarray, q: float = 0.985) -> np.ndarray: def tinted_cmap(name: str, low: str, high: str) -> LinearSegmentedColormap: - return LinearSegmentedColormap.from_list(name, [(0.0, BG), (0.24, low), (1.0, high)], N=256) + return LinearSegmentedColormap.from_list( + name, [(0.0, BG), (0.24, low), (1.0, high)], N=256 + ) def load_volume(data_dir: Path) -> tuple[np.ndarray, dict, list[str]]: @@ -59,7 +61,9 @@ def load_volume(data_dir: Path) -> tuple[np.ndarray, dict, list[str]]: return volume, summary, labels -def cube_edges(x0: float, x1: float, y0: float, y1: float, z0: float, z1: float) -> list[list[tuple[float, float, float]]]: +def cube_edges( + x0: float, x1: float, y0: float, y1: float, z0: float, z1: float +) -> list[list[tuple[float, float, float]]]: p = { "000": (x0, y0, z0), "100": (x1, y0, z0), @@ -71,9 +75,18 @@ def cube_edges(x0: float, x1: float, y0: float, y1: float, z0: float, z1: float) "111": (x1, y1, z1), } return [ - [p["000"], p["100"]], [p["010"], p["110"]], [p["001"], p["101"]], [p["011"], p["111"]], - [p["000"], p["010"]], [p["100"], p["110"]], [p["001"], p["011"]], [p["101"], p["111"]], - [p["000"], p["001"]], [p["100"], p["101"]], [p["010"], p["011"]], [p["110"], p["111"]], + [p["000"], p["100"]], + [p["010"], p["110"]], + [p["001"], p["101"]], + [p["011"], p["111"]], + [p["000"], p["010"]], + [p["100"], p["110"]], + [p["001"], p["011"]], + [p["101"], p["111"]], + [p["000"], p["001"]], + [p["100"], p["101"]], + [p["010"], p["011"]], + [p["110"], p["111"]], ] @@ -99,7 +112,11 @@ def style_3d(ax) -> None: def add_volume(ax, volume: np.ndarray) -> None: - cmaps = [tinted_cmap("answer_ct", "#063842", CYAN), tinted_cmap("random_ct", "#461813", RED), tinted_cmap("ratio_ct", "#3c2b05", AMBER)] + cmaps = [ + tinted_cmap("answer_ct", "#063842", CYAN), + tinted_cmap("random_ct", "#461813", RED), + tinted_cmap("ratio_ct", "#3c2b05", AMBER), + ] edge_colors = [CYAN, RED, AMBER] layers = np.arange(volume.shape[1]) bins = np.arange(volume.shape[2]) @@ -111,16 +128,42 @@ def add_volume(ax, volume: np.ndarray) -> None: rgba = cmap(vals) rgba[..., 3] = 0.08 + 0.68 * np.power(vals, 1.55) y = np.full_like(x, cond, dtype=np.float32) - ax.plot_surface(x, y, z, facecolors=rgba, rstride=1, cstride=1, linewidth=0, antialiased=False, shade=False) + ax.plot_surface( + x, + y, + z, + facecolors=rgba, + rstride=1, + cstride=1, + linewidth=0, + antialiased=False, + shade=False, + ) # Bright activation voxels above each condition's 98th percentile. threshold = float(np.quantile(vals, 0.982)) zz, xx = np.where(vals >= threshold) yy = np.full(xx.shape, cond, dtype=np.float32) strength = vals[zz, xx] - ax.scatter(xx, yy, zz, s=10 + 90 * strength, c=edge_colors[cond], marker="s", alpha=0.58, depthshade=False, linewidths=0) + ax.scatter( + xx, + yy, + zz, + s=10 + 90 * strength, + c=edge_colors[cond], + marker="s", + alpha=0.58, + depthshade=False, + linewidths=0, + ) - ax.add_collection3d(Line3DCollection(cube_edges(0, 179, -0.23, 2.23, 0, 18), colors=(0.42, 0.72, 0.82, 0.22), linewidths=0.9)) + ax.add_collection3d( + Line3DCollection( + cube_edges(0, 179, -0.23, 2.23, 0, 18), + colors=(0.42, 0.72, 0.82, 0.22), + linewidths=0.9, + ) + ) # Crosshair slices through the strongest answer/random separation. ratio = volume[2] @@ -128,15 +171,44 @@ def add_volume(ax, volume: np.ndarray) -> None: bin_profile = ratio.mean(axis=0) peak_layer = int(layer_profile.argmax()) peak_bin = int(bin_profile.argmax()) - ax.plot([peak_bin, peak_bin], [-0.28, 2.28], [peak_layer, peak_layer], color=GREEN, alpha=0.9, linewidth=1.5) - ax.plot([0, 179], [2.28, 2.28], [peak_layer, peak_layer], color=GREEN, alpha=0.45, linewidth=1.1) - ax.text(peak_bin + 3, 2.35, peak_layer + 0.2, "hottest ratio slice", color=GREEN, fontsize=8) + ax.plot( + [peak_bin, peak_bin], + [-0.28, 2.28], + [peak_layer, peak_layer], + color=GREEN, + alpha=0.9, + linewidth=1.5, + ) + ax.plot( + [0, 179], + [2.28, 2.28], + [peak_layer, peak_layer], + color=GREEN, + alpha=0.45, + linewidth=1.1, + ) + ax.text( + peak_bin + 3, + 2.35, + peak_layer + 0.2, + "hottest ratio slice", + color=GREEN, + fontsize=8, + ) def add_projection_panel(ax, volume: np.ndarray, labels: list[str]) -> None: ax.set_facecolor(PANEL) cmap = tinted_cmap("small_ct", "#10252f", "#f2d87b") - strip = np.vstack([volume[0], np.full((2, volume.shape[2]), np.nan), volume[1], np.full((2, volume.shape[2]), np.nan), volume[2]]) + strip = np.vstack( + [ + volume[0], + np.full((2, volume.shape[2]), np.nan), + volume[1], + np.full((2, volume.shape[2]), np.nan), + volume[2], + ] + ) masked = np.ma.masked_invalid(strip) cmap.set_bad(PANEL) ax.imshow(masked, aspect="auto", interpolation="nearest", cmap=cmap, vmin=0, vmax=1) @@ -155,7 +227,13 @@ def add_layer_panel(ax, volume: np.ndarray) -> None: names = ["answer", "random", "ratio"] for cond, color in enumerate(colors): profile = volume[cond].mean(axis=1) - ax.plot(np.arange(profile.size), profile, color=color, linewidth=2.0, label=names[cond]) + ax.plot( + np.arange(profile.size), + profile, + color=color, + linewidth=2.0, + label=names[cond], + ) ax.fill_between(np.arange(profile.size), profile, 0, color=color, alpha=0.08) ax.set_xlim(0, 18) ax.set_ylim(0, 1.0) @@ -169,21 +247,67 @@ def add_layer_panel(ax, volume: np.ndarray) -> None: spine.set_color("#27323a") -def render(volume: np.ndarray, summary: dict, labels: list[str], out_path: Path) -> None: +def render( + volume: np.ndarray, summary: dict, labels: list[str], out_path: Path +) -> None: fig = plt.figure(figsize=(18, 11), dpi=180, facecolor=BG) - gs = fig.add_gridspec(3, 5, width_ratios=[1.35, 1.35, 1.35, 0.95, 0.95], height_ratios=[0.12, 1.0, 0.42], wspace=0.22, hspace=0.24) + gs = fig.add_gridspec( + 3, + 5, + width_ratios=[1.35, 1.35, 1.35, 0.95, 0.95], + height_ratios=[0.12, 1.0, 0.42], + wspace=0.22, + hspace=0.24, + ) title_ax = fig.add_subplot(gs[0, :]) title_ax.axis("off") - title_ax.text(0.0, 0.70, "SNAPCOMPACT ACTIVATION CT", color=INK, fontsize=27, fontweight="bold", transform=title_ax.transAxes) - title_ax.text(0.0, 0.24, "volumetric tensor cube: 19 layers × 180 image-token bins × 3 conditions", color=MUTED, fontsize=11, transform=title_ax.transAxes) - title_ax.text(0.985, 0.58, f"PaddleOCR-VL · Q: {summary['question']['q']}", color=MUTED, fontsize=9, ha="right", transform=title_ax.transAxes) - title_ax.text(0.985, 0.24, f"gold answer {summary['question']['answer_text']} · answer/random mean Δ {summary['answer_over_random_delta']:.2f}×", color=AMBER, fontsize=10, ha="right", transform=title_ax.transAxes) + title_ax.text( + 0.0, + 0.70, + "SNAPCOMPACT ACTIVATION CT", + color=INK, + fontsize=27, + fontweight="bold", + transform=title_ax.transAxes, + ) + title_ax.text( + 0.0, + 0.24, + "volumetric tensor cube: 19 layers × 180 image-token bins × 3 conditions", + color=MUTED, + fontsize=11, + transform=title_ax.transAxes, + ) + title_ax.text( + 0.985, + 0.58, + f"PaddleOCR-VL · Q: {summary['question']['q']}", + color=MUTED, + fontsize=9, + ha="right", + transform=title_ax.transAxes, + ) + title_ax.text( + 0.985, + 0.24, + f"gold answer {summary['question']['answer_text']} · answer/random mean Δ {summary['answer_over_random_delta']:.2f}×", + color=AMBER, + fontsize=10, + ha="right", + transform=title_ax.transAxes, + ) ax3d = fig.add_subplot(gs[1:, :3], projection="3d") style_3d(ax3d) add_volume(ax3d, volume) - ax3d.set_title("MRI-style scan of hidden-state deltas", color=INK, fontsize=15, loc="left", pad=12) + ax3d.set_title( + "MRI-style scan of hidden-state deltas", + color=INK, + fontsize=15, + loc="left", + pad=12, + ) ax_proj = fig.add_subplot(gs[1, 3:]) add_projection_panel(ax_proj, volume, labels) @@ -191,8 +315,20 @@ def render(volume: np.ndarray, summary: dict, labels: list[str], out_path: Path) ax_layer = fig.add_subplot(gs[2, 3:]) add_layer_panel(ax_layer, volume) - fig.text(0.055, 0.055, "source: heatmaps.npz arrays answer_binned, random_binned, ratio_binned · quantile normalized per condition", color="#617078", fontsize=8) - fig.text(0.055, 0.033, "cyan=answer evidence · red=random control · gold=answer/random amplification · green=crosshair at peak ratio slice", color="#617078", fontsize=8) + fig.text( + 0.055, + 0.055, + "source: heatmaps.npz arrays answer_binned, random_binned, ratio_binned · quantile normalized per condition", + color="#617078", + fontsize=8, + ) + fig.text( + 0.055, + 0.033, + "cyan=answer evidence · red=random control · gold=answer/random amplification · green=crosshair at peak ratio slice", + color="#617078", + fontsize=8, + ) out_path.parent.mkdir(parents=True, exist_ok=True) fig.savefig(out_path, facecolor=BG, bbox_inches="tight", pad_inches=0.22) @@ -210,7 +346,11 @@ def write_source_data(volume: np.ndarray, summary: dict, out_dir: Path) -> None: ) payload = { "description": "Quantile-normalized tensor cube used by snapcompact_viz_volume.py.", - "shape": {"condition": 3, "layers": int(volume.shape[1]), "image_token_bins": int(volume.shape[2])}, + "shape": { + "condition": 3, + "layers": int(volume.shape[1]), + "image_token_bins": int(volume.shape[2]), + }, "conditions": ["answer_delta", "random_delta", "answer_over_random_ratio"], "question": summary["question"], "answer_over_random_delta": summary["answer_over_random_delta"], diff --git a/packages/snapcompact/research/snapcompact_viz_waterfall.py b/packages/snapcompact/research/snapcompact_viz_waterfall.py index de4bf49f4..e2e408885 100755 --- a/packages/snapcompact/research/snapcompact_viz_waterfall.py +++ b/packages/snapcompact/research/snapcompact_viz_waterfall.py @@ -49,17 +49,24 @@ def load_source() -> tuple[dict, dict[str, np.ndarray]]: return summary, arrays -def build_waterfall_data(summary: dict, arrays: dict[str, np.ndarray]) -> dict[str, np.ndarray]: +def build_waterfall_data( + summary: dict, arrays: dict[str, np.ndarray] +) -> dict[str, np.ndarray]: answer = smooth_rows(arrays["answer_binned"], radius=3) random = smooth_rows(arrays["random_binned"], radius=3) ratio = smooth_rows(arrays["ratio_binned"], radius=2) # Use one common robust scale so answer/random amplitudes are visually comparable. - common_high = float(summary.get("common_delta_scale_p98") or np.percentile(np.r_[answer, random], 98)) + common_high = float( + summary.get("common_delta_scale_p98") + or np.percentile(np.r_[answer, random], 98) + ) answer_u = robust_unit(answer, common_high) random_u = robust_unit(random, common_high) contrast = np.tanh((answer_u - random_u) * 2.8).astype(np.float32) - ratio_u = robust_unit(ratio, float(summary.get("ratio_scale_p98") or np.percentile(ratio, 98))) + ratio_u = robust_unit( + ratio, float(summary.get("ratio_scale_p98") or np.percentile(ratio, 98)) + ) layers, bins = answer.shape x = np.linspace(0.0, 1.0, bins, dtype=np.float32) @@ -78,7 +85,9 @@ def build_waterfall_data(summary: dict, arrays: dict[str, np.ndarray]) -> dict[s def draw_glow_line(ax, x, y, color, lw=1.4, z=5, alpha=1.0): - line, = ax.plot(x, y, color=color, lw=lw, alpha=alpha, zorder=z, solid_joinstyle="round") + (line,) = ax.plot( + x, y, color=color, lw=lw, alpha=alpha, zorder=z, solid_joinstyle="round" + ) line.set_path_effects( [ pe.Stroke(linewidth=lw + 8.5, foreground=color, alpha=0.055), @@ -238,20 +247,66 @@ def render(summary: dict, data: dict[str, np.ndarray]) -> None: # Legend built as luminous calibration bars. legend_y = -1.35 ax.plot([0.02, 0.10], [legend_y, legend_y], color=answer_color, lw=2.2) - ax.text(0.112, legend_y, "answer-mask ridge", color="#ffdca0", fontsize=9, va="center", family="monospace") + ax.text( + 0.112, + legend_y, + "answer-mask ridge", + color="#ffdca0", + fontsize=9, + va="center", + family="monospace", + ) ax.plot([0.32, 0.40], [legend_y, legend_y], color=random_color, lw=2.2) - ax.text(0.412, legend_y, "random-mask ridge", color="#9ff0ff", fontsize=9, va="center", family="monospace") + ax.text( + 0.412, + legend_y, + "random-mask ridge", + color="#9ff0ff", + fontsize=9, + va="center", + family="monospace", + ) ax.plot([0.62, 0.70], [legend_y, legend_y], color=gain_color, lw=1.4) - ax.text(0.712, legend_y, "answer excess tremor", color="#ff9bbb", fontsize=9, va="center", family="monospace") + ax.text( + 0.712, + legend_y, + "answer excess tremor", + color="#ff9bbb", + fontsize=9, + va="center", + family="monospace", + ) # Outer phosphor frame. - ax.add_patch(Rectangle((0, -0.78), 1, layers - 0.44, fill=False, lw=0.9, edgecolor="#5fb7ff", alpha=0.34, zorder=10)) + ax.add_patch( + Rectangle( + (0, -0.78), + 1, + layers - 0.44, + fill=False, + lw=0.9, + edgecolor="#5fb7ff", + alpha=0.34, + zorder=10, + ) + ) ax.set_xlim(-0.055, 1.02) ax.set_ylim(-1.62, layers + 0.98) ax.set_xticks(np.linspace(0, 1, 7)) - ax.set_xticklabels([f"{int(t * image_tokens):03d}" for t in np.linspace(0, 1, 7)], color="#8fbede", fontsize=8, family="monospace") + ax.set_xticklabels( + [f"{int(t * image_tokens):03d}" for t in np.linspace(0, 1, 7)], + color="#8fbede", + fontsize=8, + family="monospace", + ) ax.set_yticks([]) - ax.set_xlabel("image-token bin →", color="#9cc7e5", fontsize=10, family="monospace", labelpad=12) + ax.set_xlabel( + "image-token bin →", + color="#9cc7e5", + fontsize=10, + family="monospace", + labelpad=12, + ) for spine in ax.spines.values(): spine.set_visible(False) ax.tick_params(axis="x", length=0) @@ -286,7 +341,9 @@ def render(summary: dict, data: dict[str, np.ndarray]) -> None: indent=2, ) - fig.savefig(out_png, facecolor=fig.get_facecolor(), bbox_inches="tight", pad_inches=0.14) + fig.savefig( + out_png, facecolor=fig.get_facecolor(), bbox_inches="tight", pad_inches=0.14 + ) plt.close(fig) print(out_png) diff --git a/packages/snapcompact/research/squad.py b/packages/snapcompact/research/squad.py index 39d206f4d..701b20311 100644 --- a/packages/snapcompact/research/squad.py +++ b/packages/snapcompact/research/squad.py @@ -20,11 +20,19 @@ def load_paragraphs(cache: Path) -> list[dict]: out = [] for art in data: for p in art["paragraphs"]: - out.append({"ctx": " ".join(p["context"].split()), "qas": p["qas"], "title": art["title"]}) + out.append( + { + "ctx": " ".join(p["context"].split()), + "qas": p["qas"], + "title": art["title"], + } + ) return out -def build_flow(paras: list[dict], max_chars: int | None = None) -> tuple[str, list[int]]: +def build_flow( + paras: list[dict], max_chars: int | None = None +) -> tuple[str, list[int]]: """Space-joined passage stream + start offset of each passage.""" flow, offsets = "", [] for p in paras: @@ -95,7 +103,7 @@ def f1(pred: str, golds: list[str]) -> float: def parse_numbered(text: str, n: int) -> list[str]: - """Extract answers from a numbered list; missing entries become ''. """ + """Extract answers from a numbered list; missing entries become ''.""" answers = [""] * n for line in text.splitlines(): m = re.match(r"\s*(\d+)[.):]\s*(.*\S)?\s*$", line) diff --git a/python/omp-rpc/src/omp_rpc/protocol.py b/python/omp-rpc/src/omp_rpc/protocol.py index ccd94cee7..5cc4ad3c6 100644 --- a/python/omp-rpc/src/omp_rpc/protocol.py +++ b/python/omp-rpc/src/omp_rpc/protocol.py @@ -1144,7 +1144,10 @@ def _parse_thinking_config(payload: object) -> ThinkingConfig | None: if not isinstance(raw_efforts, list): raise ValueError("model.thinking.efforts must be a list") efforts: tuple[Effort, ...] = tuple( - cast(Effort, _require_literal(item, _EFFORT_VALUES, field="model.thinking.efforts[]")) + cast( + Effort, + _require_literal(item, _EFFORT_VALUES, field="model.thinking.efforts[]"), + ) for item in raw_efforts ) return ThinkingConfig( diff --git a/scripts/analyze_small_edits.py b/scripts/analyze_small_edits.py index 616182666..4b571fcbe 100755 --- a/scripts/analyze_small_edits.py +++ b/scripts/analyze_small_edits.py @@ -12,9 +12,21 @@ import sys if __package__ in (None, ""): sys.path.insert(0, str(Path(__file__).resolve().parent)) - from tool_io import ReservoirSample, ToolIOConfig, ToolInvocation, iter_tool_invocations, list_recent_session_files + from tool_io import ( + ReservoirSample, + ToolIOConfig, + ToolInvocation, + iter_tool_invocations, + list_recent_session_files, + ) else: - from scripts.tool_io import ReservoirSample, ToolIOConfig, ToolInvocation, iter_tool_invocations, list_recent_session_files + from scripts.tool_io import ( + ReservoirSample, + ToolIOConfig, + ToolInvocation, + iter_tool_invocations, + list_recent_session_files, + ) TOOL_NAMES = ("edit", "ast_edit") @@ -81,10 +93,13 @@ class RunStats: small_edits_after_same_path_failed_edit: int = 0 - def parse_args() -> argparse.Namespace: - parser = argparse.ArgumentParser(description="Analyze small edit/ast_edit tool usage in session logs.") - parser.add_argument("--sessions-dir", type=Path, default=Path.home() / ".omp" / "agent" / "sessions") + parser = argparse.ArgumentParser( + description="Analyze small edit/ast_edit tool usage in session logs." + ) + parser.add_argument( + "--sessions-dir", type=Path, default=Path.home() / ".omp" / "agent" / "sessions" + ) parser.add_argument("--sample-size", type=positive_int, default=30) parser.add_argument("--max-files", type=positive_int, default=500) parser.add_argument("--since-days", type=positive_int, default=30) @@ -94,7 +109,6 @@ def parse_args() -> argparse.Namespace: return parser.parse_args() - def positive_int(value: str) -> int: parsed = int(value) if parsed <= 0: @@ -102,7 +116,6 @@ def positive_int(value: str) -> int: return parsed - def strip_decorations(line: str) -> str: return re.sub(r"^\s*\d+\s+", "", line).strip() @@ -111,7 +124,6 @@ def is_delimiter_line(line: str) -> bool: return bool(re.match(r"^[\]}),;]+$", line)) - def is_tiny_structural_line(line: str) -> bool: if len(line) == 0: return True @@ -126,14 +138,16 @@ def is_tiny_structural_line(line: str) -> bool: return False - def classify_success_issue(summary: DiffSummary) -> str: previews = summary.changed_preview if previews and all(len(line) == 0 for line in previews): return "blank-line-adjustment" if previews and all(is_delimiter_line(line) for line in previews): return "delimiter-adjustment" - if previews and all(re.match(r"^(pub\s+mod|pub\s+use|mod|use|import|export)\b", line) for line in previews): + if previews and all( + re.match(r"^(pub\s+mod|pub\s+use|mod|use|import|export)\b", line) + for line in previews + ): return "import-or-module-tweak" if summary.removed_lines == 1 and summary.added_lines == 0: return "single-line-delete" @@ -144,24 +158,32 @@ def classify_success_issue(summary: DiffSummary) -> str: return "small-structural-fix" - def classify_failure_issue(result_text: str) -> str: if re.search(r"identical content|No changes made", result_text, re.IGNORECASE): return "no-op-identical" - if re.search(r"Failed to find context|matches for context|expected lines|tag mismatch|>>>", result_text, re.IGNORECASE): + if re.search( + r"Failed to find context|matches for context|expected lines|tag mismatch|>>>", + result_text, + re.IGNORECASE, + ): return "context-mismatch" - if re.search(r"Unexpected line in hunk|parse error|SyntaxError", result_text, re.IGNORECASE): + if re.search( + r"Unexpected line in hunk|parse error|SyntaxError", result_text, re.IGNORECASE + ): return "invalid-patch-shape" if re.search(r"File not found", result_text, re.IGNORECASE): return "missing-file" if re.search(r"occurrence|ambiguous", result_text, re.IGNORECASE): return "ambiguous-target" - if re.search(r"Validation failed|required property|must have required property", result_text, re.IGNORECASE): + if re.search( + r"Validation failed|required property|must have required property", + result_text, + re.IGNORECASE, + ): return "invalid-arguments" return "other-failure" - def summarize_diff(diff: str | None) -> DiffSummary: if not diff: return DiffSummary( @@ -188,7 +210,9 @@ def summarize_diff(diff: str | None) -> DiffSummary: previews = [line for line in all_changes if line or include_blank] changed_lines = len(added) + len(removed) tiny_only = all(is_tiny_structural_line(line) for line in all_changes) - small = changed_lines > 0 and (changed_lines <= 2 or (changed_lines <= 4 and tiny_only)) + small = changed_lines > 0 and ( + changed_lines <= 2 or (changed_lines <= 4 and tiny_only) + ) preview_slice = previews[:4] category = None if small: @@ -211,18 +235,20 @@ def summarize_diff(diff: str | None) -> DiffSummary: ) - def previews_need_blank_marker(lines: list[str]) -> bool: return any(len(line) == 0 for line in lines) - def build_completed_edit(invocation: ToolInvocation) -> CompletedEdit | None: if not invocation.has_result: return None diff_summary = summarize_diff(invocation.diff) is_error = invocation.is_error - issue = classify_failure_issue(invocation.result_text) if is_error else (diff_summary.category or "other-success") + issue = ( + classify_failure_issue(invocation.result_text) + if is_error + else (diff_summary.category or "other-success") + ) return CompletedEdit( session_file=str(invocation.session_file), tool_call_id=invocation.tool_call_id, @@ -245,8 +271,9 @@ def build_completed_edit(invocation: ToolInvocation) -> CompletedEdit | None: ) - -def analyze_small_edits(stream: Iterable[ToolInvocation], *, sample_size: int, files_scanned: int) -> dict[str, object]: +def analyze_small_edits( + stream: Iterable[ToolInvocation], *, sample_size: int, files_scanned: int +) -> dict[str, object]: sample: ReservoirSample[Candidate] = ReservoirSample(size=sample_size) issue_counts: dict[str, int] = {} stats = RunStats(files_scanned=files_scanned) @@ -295,7 +322,6 @@ def analyze_small_edits(stream: Iterable[ToolInvocation], *, sample_size: int, f } - def candidate_to_dict(candidate: Candidate) -> dict[str, object]: payload = {"kind": candidate.kind, "edit": asdict(candidate.edit)} if candidate.previous_edit is not None: @@ -303,19 +329,20 @@ def candidate_to_dict(candidate: Candidate) -> dict[str, object]: return payload - def top_entries(counts: dict[str, int], limit: int) -> list[dict[str, object]]: return [ {"name": name, "count": count} - for name, count in sorted(counts.items(), key=lambda entry: (-entry[1], entry[0]))[:limit] + for name, count in sorted( + counts.items(), key=lambda entry: (-entry[1], entry[0]) + )[:limit] ] - def short_path(target_path: str) -> str: home = str(Path.home()) - return f"~{target_path[len(home):]}" if target_path.startswith(home) else target_path - + return ( + f"~{target_path[len(home) :]}" if target_path.startswith(home) else target_path + ) def truncate(text: str, limit: int) -> str: @@ -324,7 +351,6 @@ def truncate(text: str, limit: int) -> str: return f"{text[: limit - 1]}…" - def format_sample_entry(candidate: dict[str, object], index: int) -> str: edit = candidate["edit"] assert isinstance(edit, dict) @@ -340,10 +366,16 @@ def format_sample_entry(candidate: dict[str, object], index: int) -> str: ) changed_preview = edit.get("changed_preview") if isinstance(changed_preview, list) and changed_preview: - lines.append(f" preview: {' | '.join(str(item) for item in changed_preview)}") + lines.append( + f" preview: {' | '.join(str(item) for item in changed_preview)}" + ) previous = candidate.get("previous_edit") if isinstance(previous, dict): - path_part = f" ({short_path(str(previous['path']))})" if previous.get("path") else "" + path_part = ( + f" ({short_path(str(previous['path']))})" + if previous.get("path") + else "" + ) lines.append( " previous edit: " f"{'same-path' if previous.get('same_path') else 'other-path'} " @@ -351,16 +383,21 @@ def format_sample_entry(candidate: dict[str, object], index: int) -> str: ) previous_preview = previous.get("changed_preview") if isinstance(previous_preview, list) and previous_preview: - lines.append(f" previous preview: {' | '.join(str(item) for item in previous_preview)}") + lines.append( + f" previous preview: {' | '.join(str(item) for item in previous_preview)}" + ) else: lines.append(" previous edit: none") else: - lines.append(f" result: {truncate(' '.join(str(edit['result_text']).split()), 220)}") + lines.append( + f" result: {truncate(' '.join(str(edit['result_text']).split()), 220)}" + ) changed_preview = edit.get("changed_preview") if isinstance(changed_preview, list) and changed_preview: - lines.append(f" diff preview: {' | '.join(str(item) for item in changed_preview)}") - return '\n'.join(lines) - + lines.append( + f" diff preview: {' | '.join(str(item) for item in changed_preview)}" + ) + return "\n".join(lines) def main() -> None: @@ -375,7 +412,9 @@ def main() -> None: ) files = list_recent_session_files(config) stream = iter_tool_invocations(TOOL_NAMES, config) - analysis = analyze_small_edits(stream, sample_size=options.sample_size, files_scanned=len(files)) + analysis = analyze_small_edits( + stream, sample_size=options.sample_size, files_scanned=len(files) + ) if options.json: print( @@ -403,14 +442,20 @@ def main() -> None: assert isinstance(stats, dict) assert isinstance(top_issues, list) assert isinstance(sample, list) - print(f"Scanned {stats['files_scanned']} session file(s) from {short_path(str(options.sessions_dir))}") + print( + f"Scanned {stats['files_scanned']} session file(s) from {short_path(str(options.sessions_dir))}" + ) print(f"Edit attempts: {stats['total_edit_attempts']}") print(f"Failed edits: {stats['failed_edits']}") print(f"Small edits: {stats['small_edits']}") print(f"Small edits with previous edit: {stats['small_edits_with_previous_edit']}") - print(f"Small edits with previous same-path edit: {stats['small_edits_with_previous_same_path']}") + print( + f"Small edits with previous same-path edit: {stats['small_edits_with_previous_same_path']}" + ) print(f"Small edits after failed edit: {stats['small_edits_after_failed_edit']}") - print(f"Small edits after same-path failed edit: {stats['small_edits_after_same_path_failed_edit']}") + print( + f"Small edits after same-path failed edit: {stats['small_edits_after_same_path_failed_edit']}" + ) print() print("Top issues:") for entry in top_issues[:12]: diff --git a/scripts/edit-benchmark.py b/scripts/edit-benchmark.py index 86d1926dc..be4c01acd 100755 --- a/scripts/edit-benchmark.py +++ b/scripts/edit-benchmark.py @@ -6,68 +6,77 @@ Select the edit variant via the PI_EDIT_VARIANT env var (e.g. `vim`, `hashline`, `replace`, `patch`, `apply_patch`) or `--variant`. Examples: - PI_EDIT_VARIANT=vim scripts/edit-benchmark.py - scripts/edit-benchmark.py --variant hashline + PI_EDIT_VARIANT=vim scripts/edit-benchmark.py + scripts/edit-benchmark.py --variant hashline """ + from __future__ import annotations import os import sys -from edit_benchmark_common import BenchmarkSpec, EDIT_DIFF, EXPECTED_CONTENT, run_benchmark_main +from edit_benchmark_common import ( + BenchmarkSpec, + EDIT_DIFF, + EXPECTED_CONTENT, + run_benchmark_main, +) + def _extract_variant_arg() -> str | None: - """Pop `--variant ` (or `--variant=`) from sys.argv before argparse in common runs.""" - argv = sys.argv - for i, arg in enumerate(argv[1:], start=1): - if arg == "--variant" and i + 1 < len(argv): - value = argv[i + 1] - del argv[i : i + 2] - return value - if arg.startswith("--variant="): - value = arg.split("=", 1)[1] - del argv[i] - return value - return None + """Pop `--variant ` (or `--variant=`) from sys.argv before argparse in common runs.""" + argv = sys.argv + for i, arg in enumerate(argv[1:], start=1): + if arg == "--variant" and i + 1 < len(argv): + value = argv[i + 1] + del argv[i : i + 2] + return value + if arg.startswith("--variant="): + value = arg.split("=", 1)[1] + del argv[i] + return value + return None def _resolve_variant() -> str: - cli_variant = _extract_variant_arg() - variant = cli_variant or os.environ.get("PI_EDIT_VARIANT") - if not variant: - raise SystemExit("edit-benchmark: set PI_EDIT_VARIANT= or pass --variant .") - return variant + cli_variant = _extract_variant_arg() + variant = cli_variant or os.environ.get("PI_EDIT_VARIANT") + if not variant: + raise SystemExit( + "edit-benchmark: set PI_EDIT_VARIANT= or pass --variant ." + ) + return variant def build_spec(variant: str) -> BenchmarkSpec: - mode_phrase = f"in {variant} mode" - prompt = ( - f"Use the `read` tool to inspect `test.rs`, then use the `edit` tool {mode_phrase} " - f"to make `test.rs` exactly match the requested change.\n" - f"\n" - f"Apply this diff:\n" - f"```diff\n" - f"{EDIT_DIFF}```\n" - f"\n" - f"Final expected file content:\n" - f"```rust\n" - f"{EXPECTED_CONTENT}```\n" - ) - retry = f"Please try again using the edit tool {mode_phrase}." - return BenchmarkSpec( - description=f"Benchmark edit tool in {variant} mode across models with simple edit tasks.", - workspace_prefix=f"{variant}-benchmark", - tools=("edit", "read"), - env={"PI_EDIT_VARIANT": variant, "PI_STRICT_EDIT_MODE": "1"}, - initial_prompt=prompt, - retry_instruction=retry, - ) + mode_phrase = f"in {variant} mode" + prompt = ( + f"Use the `read` tool to inspect `test.rs`, then use the `edit` tool {mode_phrase} " + f"to make `test.rs` exactly match the requested change.\n" + f"\n" + f"Apply this diff:\n" + f"```diff\n" + f"{EDIT_DIFF}```\n" + f"\n" + f"Final expected file content:\n" + f"```rust\n" + f"{EXPECTED_CONTENT}```\n" + ) + retry = f"Please try again using the edit tool {mode_phrase}." + return BenchmarkSpec( + description=f"Benchmark edit tool in {variant} mode across models with simple edit tasks.", + workspace_prefix=f"{variant}-benchmark", + tools=("edit", "read"), + env={"PI_EDIT_VARIANT": variant, "PI_STRICT_EDIT_MODE": "1"}, + initial_prompt=prompt, + retry_instruction=retry, + ) def main() -> int: - variant = _resolve_variant() - return run_benchmark_main(build_spec(variant)) + variant = _resolve_variant() + return run_benchmark_main(build_spec(variant)) if __name__ == "__main__": - raise SystemExit(main()) + raise SystemExit(main()) diff --git a/scripts/edit_benchmark_common.py b/scripts/edit_benchmark_common.py index 3280fc1a8..edb952d8e 100644 --- a/scripts/edit_benchmark_common.py +++ b/scripts/edit_benchmark_common.py @@ -2,6 +2,7 @@ """ Shared helpers for edit benchmark scripts. """ + from __future__ import annotations import argparse @@ -21,13 +22,19 @@ from typing import Any, Callable REPO_ROOT = Path(__file__).resolve().parent.parent sys.path.insert(0, str(REPO_ROOT / "python/omp-rpc/src")) -from omp_rpc import MessageEndEvent, MessageStartEvent, MessageUpdateEvent, RpcClient, ToolExecutionStartEvent # noqa: E402 +from omp_rpc import ( + MessageEndEvent, + MessageStartEvent, + MessageUpdateEvent, + RpcClient, + ToolExecutionStartEvent, +) # noqa: E402 MODELS = [ "openrouter/moonshotai/kimi-k2.5", "openrouter/anthropic/claude-haiku-4.5", "openrouter/google/gemini-3.1-flash-lite-preview", - "openrouter/z-ai/glm-4.7-20251222:nitro" + "openrouter/z-ai/glm-4.7-20251222:nitro", # "openrouter/anthropic/claude-sonnet-4.6", # "openrouter/google/gemini-3-flash-preview", # "openrouter/z-ai/glm-5-turbo", @@ -457,6 +464,7 @@ mod tests { } """ + def _compute_edit_diff() -> str: initial_lines = INITIAL_CONTENT.splitlines(keepends=True) expected_lines = EXPECTED_CONTENT.splitlines(keepends=True) @@ -526,13 +534,17 @@ class VerbosePrinter: sys.stderr.flush() self._open_kind = None - def emit_delta(self, kind: str, delta: str, content_index: int | None = None) -> None: + def emit_delta( + self, kind: str, delta: str, content_index: int | None = None + ) -> None: if not delta: return if content_index is not None: key = (kind, content_index) - self._seen_block_lengths[key] = self._seen_block_lengths.get(key, 0) + len(delta) + self._seen_block_lengths[key] = self._seen_block_lengths.get(key, 0) + len( + delta + ) with _PRINT_LOCK: if self._open_kind != kind: @@ -553,7 +565,9 @@ class VerbosePrinter: sys.stderr.flush() def emit_tool_call(self, tool_name: str, args: Any) -> None: - rendered_args = json.dumps(args, ensure_ascii=False, sort_keys=True, separators=(",", ":")) + rendered_args = json.dumps( + args, ensure_ascii=False, sort_keys=True, separators=(",", ":") + ) with _PRINT_LOCK: if self._open_kind is not None: sys.stderr.write("\n") @@ -597,7 +611,10 @@ class VerbosePrinter: if not isinstance(content, list): return - has_redacted = any(isinstance(block, dict) and block.get("type") == "redactedThinking" for block in content) + has_redacted = any( + isinstance(block, dict) and block.get("type") == "redactedThinking" + for block in content + ) if not has_redacted: return @@ -624,7 +641,9 @@ def resolve_omp_bin(raw: str | None) -> str: return repo_bin found = shutil.which("omp") if not found: - raise SystemExit("Could not find `omp` on PATH and could not resolve the repo CLI. Set --omp-bin or OMP_BIN.") + raise SystemExit( + "Could not find `omp` on PATH and could not resolve the repo CLI. Set --omp-bin or OMP_BIN." + ) return found @@ -672,9 +691,13 @@ def install_verbose_logging( message_event = event.assistant_message_event event_type = message_event["type"] if event_type == "text_delta": - printer.emit_delta("text", message_event["delta"], message_event["contentIndex"]) + printer.emit_delta( + "text", message_event["delta"], message_event["contentIndex"] + ) elif event_type == "thinking_delta": - printer.emit_delta("thinking", message_event["delta"], message_event["contentIndex"]) + printer.emit_delta( + "thinking", message_event["delta"], message_event["contentIndex"] + ) def handle_message_end(event: MessageEndEvent) -> None: if not include_messages: @@ -745,7 +768,6 @@ def run_benchmark_for_model( client.install_headless_ui() verbose_cleanup = install_verbose_logging(client, model, log_mode, thinking) - def handle_tool_count(event: ToolExecutionStartEvent) -> None: nonlocal edit_tool_calls, turns_used if counting_edit_turns: @@ -816,11 +838,15 @@ def run_benchmark_for_model( ) -async def run_all(spec: BenchmarkSpec, args: argparse.Namespace) -> dict[str, dict[str, Any]]: +async def run_all( + spec: BenchmarkSpec, args: argparse.Namespace +) -> dict[str, dict[str, Any]]: omp_bin = resolve_omp_bin(args.omp_bin) timestamp = time.strftime("%Y%m%d-%H%M%S") - workspace_root = Path(tempfile.gettempdir()) / f"{spec.workspace_prefix}-{timestamp}" + workspace_root = ( + Path(tempfile.gettempdir()) / f"{spec.workspace_prefix}-{timestamp}" + ) workspace_root.mkdir(parents=True, exist_ok=True) selected_models = args.models or MODELS @@ -839,7 +865,9 @@ async def run_all(spec: BenchmarkSpec, args: argparse.Namespace) -> dict[str, di omp_bin=omp_bin, workspace=workspace, timeout=args.timeout, - log_mode="verbose" if args.verbose else ("print" if args.print else None), + log_mode="verbose" + if args.verbose + else ("print" if args.print else None), thinking=args.thinking, max_turns=args.max_turns, ) diff --git a/scripts/rate-edit-tool.py b/scripts/rate-edit-tool.py index b06027ffe..cc8aa731a 100755 --- a/scripts/rate-edit-tool.py +++ b/scripts/rate-edit-tool.py @@ -48,7 +48,7 @@ MODELS = [ "openrouter/moonshotai/kimi-k2.5", "openrouter/anthropic/claude-haiku-4.5", "openrouter/z-ai/glm-4.7", - "openai-codex/gpt-5.4" + "openai-codex/gpt-5.4", ] ORACLE_MODEL = "openrouter/anthropic/claude-opus-4.6" @@ -1472,14 +1472,18 @@ async def run_all(args: argparse.Namespace) -> int: ) except (RpcError, RpcProcessExitError) as exc: err = f"{type(exc).__name__}: {exc}" - (results_dir / "oracle_error.txt").write_text(err + "\n", encoding="utf-8") + (results_dir / "oracle_error.txt").write_text( + err + "\n", encoding="utf-8" + ) print(f"Oracle synthesis FAILED: {err}", file=sys.stderr) print(f"Saved error to {results_dir}/oracle_error.txt", file=sys.stderr) return 2 print(synthesis) return 0 combined = format_combined_reviews(sources) - (results_dir / "combined_reviews.md").write_text(combined + "\n", encoding="utf-8") + (results_dir / "combined_reviews.md").write_text( + combined + "\n", encoding="utf-8" + ) print(combined) return 0 diff --git a/scripts/session-stats/analyze_search_relevance.py b/scripts/session-stats/analyze_search_relevance.py index b84bd6c6c..2cae39fb5 100644 --- a/scripts/session-stats/analyze_search_relevance.py +++ b/scripts/session-stats/analyze_search_relevance.py @@ -19,6 +19,7 @@ whichever comes first. Outputs scripts/session-stats/out/search-relevance.png. """ + from __future__ import annotations import argparse @@ -36,8 +37,8 @@ import numpy as np DB_PATH = Path.home() / ".omp" / "stats.db" OUT_DIR = Path(__file__).resolve().parent / "out" -DEFAULT_SINCE = "2026-04-01" # search/grep traffic before this is sparse -LOOKAHEAD = 30 # max tool calls to scan after a search +DEFAULT_SINCE = "2026-04-01" # search/grep traffic before this is sparse +LOOKAHEAD = 30 # max tool calls to scan after a search # --------------------------------------------------------------------------- # @@ -107,6 +108,7 @@ def extract_paths(result_text: str | None) -> tuple[list[str], dict[str, int]]: # --------------------------------------------------------------------------- # # Tool-call helpers + def search_signature(arg_obj: dict) -> tuple: """Stable identity key for a search/grep call: (pattern, path-scope). @@ -154,6 +156,7 @@ def read_path(arg_obj: dict) -> str | None: # --------------------------------------------------------------------------- # # Per-session walk + def classify_sessions(conn: sqlite3.Connection, since_ms: int) -> list[dict]: """Walks each session in seq order, classifying every search/grep call.""" # Pull calls + paired results in one ordered stream per session. @@ -301,23 +304,33 @@ def report(records: list[dict]) -> None: coverage = deepest_1b / result_counts engaged_n = np.array([r["engaged_count"] for r in engaged], dtype=np.int64) print(f"\nfor engaged-read calls (n={len(engaged):,}):") - print(f" deepest index reached p50={int(np.median(deepest_1b))} " - f"p75={int(np.percentile(deepest_1b,75))} " - f"p90={int(np.percentile(deepest_1b,90))} " - f"max={int(deepest_1b.max())}") - print(f" result list length p50={int(np.median(result_counts))} " - f"p90={int(np.percentile(result_counts,90))} " - f"max={int(result_counts.max())}") - print(f" deepest / list size p50={np.median(coverage)*100:.0f}% " - f"p25={np.percentile(coverage,25)*100:.0f}%") - print(f" reads per result list p50={int(np.median(engaged_n))} " - f"p90={int(np.percentile(engaged_n,90))}") + print( + f" deepest index reached p50={int(np.median(deepest_1b))} " + f"p75={int(np.percentile(deepest_1b, 75))} " + f"p90={int(np.percentile(deepest_1b, 90))} " + f"max={int(deepest_1b.max())}" + ) + print( + f" result list length p50={int(np.median(result_counts))} " + f"p90={int(np.percentile(result_counts, 90))} " + f"max={int(result_counts.max())}" + ) + print( + f" deepest / list size p50={np.median(coverage) * 100:.0f}% " + f"p25={np.percentile(coverage, 25) * 100:.0f}%" + ) + print( + f" reads per result list p50={int(np.median(engaged_n))} " + f"p90={int(np.percentile(engaged_n, 90))}" + ) next_page = sum(1 for r in records if r["next_page"]) refined = sum(1 for r in records if r["refined"]) print(f"\nbehaviours (not exclusive):") - print(f" any next-page request : {next_page:,} ({100*next_page/total:.1f}%)") - print(f" any refined-query : {refined:,} ({100*refined/total:.1f}%)") + print( + f" any next-page request : {next_page:,} ({100 * next_page / total:.1f}%)" + ) + print(f" any refined-query : {refined:,} ({100 * refined / total:.1f}%)") # Shape of result lists — files per result, matches per file, and whether # diversity (files-per-result / matches-per-file) correlates with engagement. @@ -327,26 +340,39 @@ def report(records: list[dict]) -> None: dtype=np.int64, ) print(f"\nresult shape across all {total:,} calls:") - print(f" files per result " - f"p50={int(np.median(files_per_result))} " - f"p75={int(np.percentile(files_per_result,75))} " - f"p90={int(np.percentile(files_per_result,90))} " - f"p99={int(np.percentile(files_per_result,99))} " - f"max={int(files_per_result.max())}") + print( + f" files per result " + f"p50={int(np.median(files_per_result))} " + f"p75={int(np.percentile(files_per_result, 75))} " + f"p90={int(np.percentile(files_per_result, 90))} " + f"p99={int(np.percentile(files_per_result, 99))} " + f"max={int(files_per_result.max())}" + ) if matches_per_file_flat.size: - print(f" matches per file (flat) " - f"p50={int(np.median(matches_per_file_flat))} " - f"p75={int(np.percentile(matches_per_file_flat,75))} " - f"p90={int(np.percentile(matches_per_file_flat,90))} " - f"p99={int(np.percentile(matches_per_file_flat,99))} " - f"max={int(matches_per_file_flat.max())}") + print( + f" matches per file (flat) " + f"p50={int(np.median(matches_per_file_flat))} " + f"p75={int(np.percentile(matches_per_file_flat, 75))} " + f"p90={int(np.percentile(matches_per_file_flat, 90))} " + f"p99={int(np.percentile(matches_per_file_flat, 99))} " + f"max={int(matches_per_file_flat.max())}" + ) # Engagement vs shape: is the model more likely to read at all when there # are more distinct files? When matches are more concentrated per file? print(f"\nengagement vs result shape:") - print(f" {'files-per-result':<22} {'n calls':>9} {'engaged %':>10} {'p50 deepest':>12}") - bins = [(1, 1, "1"), (2, 2, "2"), (3, 5, "3-5"), (6, 10, "6-10"), - (11, 20, "11-20"), (21, 50, "21-50"), (51, 10**9, "51+")] + print( + f" {'files-per-result':<22} {'n calls':>9} {'engaged %':>10} {'p50 deepest':>12}" + ) + bins = [ + (1, 1, "1"), + (2, 2, "2"), + (3, 5, "3-5"), + (6, 10, "6-10"), + (11, 20, "11-20"), + (21, 50, "21-50"), + (51, 10**9, "51+"), + ] for lo, hi, label in bins: bucket = [r for r in records if lo <= r["n_results"] <= hi] if not bucket: @@ -359,13 +385,22 @@ def report(records: list[dict]) -> None: p50_deep = 0 print(f" {label:<22} {len(bucket):>9,} {eng_share:>9.1f}% {p50_deep:>12}") - print(f"\n {'max matches/file':<22} {'n calls':>9} {'engaged %':>10} {'p50 deepest':>12}") - bins = [(1, 1, "1"), (2, 5, "2-5"), (6, 20, "6-20"), - (21, 100, "21-100"), (101, 10**9, "100+")] + print( + f"\n {'max matches/file':<22} {'n calls':>9} {'engaged %':>10} {'p50 deepest':>12}" + ) + bins = [ + (1, 1, "1"), + (2, 5, "2-5"), + (6, 20, "6-20"), + (21, 100, "21-100"), + (101, 10**9, "100+"), + ] for lo, hi, label in bins: - bucket = [r for r in records - if r["matches_per_file"] - and lo <= max(r["matches_per_file"]) <= hi] + bucket = [ + r + for r in records + if r["matches_per_file"] and lo <= max(r["matches_per_file"]) <= hi + ] if not bucket: continue eng = [r for r in bucket if r["outcome"] == "engaged-read"] @@ -380,6 +415,7 @@ def report(records: list[dict]) -> None: # --------------------------------------------------------------------------- # # Plot + def plot(records: list[dict], since: str) -> Path | None: if not records: return None @@ -400,8 +436,14 @@ def plot(records: list[dict], since: str) -> Path | None: colors = [OUTCOME_COLORS[o] for o in ordered] bars = ax.bar(ordered, pct, color=colors, edgecolor="#1f2937", linewidth=0.5) for b, p, n in zip(bars, pct, nvals): - ax.text(b.get_x() + b.get_width() / 2, p + 1.5, - f"{p:.1f}%\nn={n:,}", ha="center", va="bottom", fontsize=9) + ax.text( + b.get_x() + b.get_width() / 2, + p + 1.5, + f"{p:.1f}%\nn={n:,}", + ha="center", + va="bottom", + fontsize=9, + ) ax.set_title(f"search outcome (n={total:,})") ax.set_ylabel("share of calls") ax.set_ylim(0, max(pct) + 12) @@ -419,8 +461,14 @@ def plot(records: list[dict], since: str) -> Path | None: pct = 100 * hist / deepest.size bars = ax.bar(labels, pct, color="#0f766e", edgecolor="#134e4a", linewidth=0.5) for b, p, n in zip(bars, pct, hist): - ax.text(b.get_x() + b.get_width() / 2, p + 1.2, - f"{p:.1f}%\nn={n:,}", ha="center", va="bottom", fontsize=8) + ax.text( + b.get_x() + b.get_width() / 2, + p + 1.2, + f"{p:.1f}%\nn={n:,}", + ha="center", + va="bottom", + fontsize=8, + ) ax.set_title("deepest result index the model read") ax.set_ylabel("share of engaged-read calls") ax.set_ylim(0, max(pct) + 12) @@ -437,8 +485,12 @@ def plot(records: list[dict], since: str) -> Path | None: coverage.sort() cdf = np.arange(1, coverage.size + 1) / coverage.size ax.plot(coverage, cdf, color="#7c3aed", linewidth=2.0) - ax.axvline(0.1, color="#9ca3af", linestyle="--", linewidth=1, label="10% of list") - ax.axvline(0.5, color="#9ca3af", linestyle=":", linewidth=1, label="50% of list") + ax.axvline( + 0.1, color="#9ca3af", linestyle="--", linewidth=1, label="10% of list" + ) + ax.axvline( + 0.5, color="#9ca3af", linestyle=":", linewidth=1, label="50% of list" + ) ax.set_title("coverage CDF — deepest read / list size") ax.set_xlabel("fraction of list reached") ax.set_ylabel("CDF of engaged-read calls") @@ -454,9 +506,14 @@ def plot(records: list[dict], since: str) -> Path | None: sizes = [r["n_results"] for r in records if r["outcome"] == outcome] if not sizes: continue - ax.hist(sizes, bins=bins, histtype="step", linewidth=1.8, - color=OUTCOME_COLORS[outcome], - label=f"{outcome} (p50={int(np.median(sizes))})") + ax.hist( + sizes, + bins=bins, + histtype="step", + linewidth=1.8, + color=OUTCOME_COLORS[outcome], + label=f"{outcome} (p50={int(np.median(sizes))})", + ) ax.set_xscale("log") ax.set_yscale("log") ax.set_xlabel("result list size") @@ -467,8 +524,15 @@ def plot(records: list[dict], since: str) -> Path | None: # Panel E — engagement rate vs files-per-result, with p50 deepest overlay. ax = axes[2, 0] - bins = [(1, 1, "1"), (2, 2, "2"), (3, 5, "3-5"), (6, 10, "6-10"), - (11, 20, "11-20"), (21, 50, "21-50"), (51, 10**9, "51+")] + bins = [ + (1, 1, "1"), + (2, 2, "2"), + (3, 5, "3-5"), + (6, 10, "6-10"), + (11, 20, "11-20"), + (21, 50, "21-50"), + (51, 10**9, "51+"), + ] labels = [] eng_share = [] deep_p50 = [] @@ -481,13 +545,27 @@ def plot(records: list[dict], since: str) -> Path | None: n_calls.append(len(bucket)) engs = [r for r in bucket if r["outcome"] == "engaged-read"] eng_share.append(100 * len(engs) / len(bucket)) - deep_p50.append(int(np.median([r["deepest_index"] + 1 for r in engs])) if engs else 0) + deep_p50.append( + int(np.median([r["deepest_index"] + 1 for r in engs])) if engs else 0 + ) x = np.arange(len(labels)) - bars = ax.bar(x, eng_share, color="#16a34a", edgecolor="#14532d", - linewidth=0.5, label="engaged %") + bars = ax.bar( + x, + eng_share, + color="#16a34a", + edgecolor="#14532d", + linewidth=0.5, + label="engaged %", + ) for b, p, n in zip(bars, eng_share, n_calls): - ax.text(b.get_x() + b.get_width() / 2, p + 0.8, - f"{p:.0f}%\nn={n:,}", ha="center", va="bottom", fontsize=8) + ax.text( + b.get_x() + b.get_width() / 2, + p + 0.8, + f"{p:.0f}%\nn={n:,}", + ha="center", + va="bottom", + fontsize=8, + ) ax.set_xticks(x) ax.set_xticklabels(labels) ax.set_ylabel("engaged %", color="#15803d") @@ -496,37 +574,64 @@ def plot(records: list[dict], since: str) -> Path | None: ax.set_title("engagement vs files-per-result") ax.set_xlabel("files in result") ax2 = ax.twinx() - ax2.plot(x, deep_p50, color="#7c3aed", marker="o", linewidth=1.8, - label="p50 deepest index") + ax2.plot( + x, + deep_p50, + color="#7c3aed", + marker="o", + linewidth=1.8, + label="p50 deepest index", + ) ax2.set_ylabel("p50 deepest index", color="#5b21b6") ax2.tick_params(axis="y", labelcolor="#5b21b6") ax.grid(True, axis="y", alpha=0.25, linestyle="--") # Panel F — engagement rate vs max matches-per-file. ax = axes[2, 1] - bins = [(1, 1, "1"), (2, 5, "2-5"), (6, 20, "6-20"), - (21, 100, "21-100"), (101, 10**9, "100+")] + bins = [ + (1, 1, "1"), + (2, 5, "2-5"), + (6, 20, "6-20"), + (21, 100, "21-100"), + (101, 10**9, "100+"), + ] labels = [] eng_share = [] deep_p50 = [] n_calls = [] for lo, hi, label in bins: - bucket = [r for r in records - if r["matches_per_file"] - and lo <= max(r["matches_per_file"]) <= hi] + bucket = [ + r + for r in records + if r["matches_per_file"] and lo <= max(r["matches_per_file"]) <= hi + ] if not bucket: continue labels.append(label) n_calls.append(len(bucket)) engs = [r for r in bucket if r["outcome"] == "engaged-read"] eng_share.append(100 * len(engs) / len(bucket)) - deep_p50.append(int(np.median([r["deepest_index"] + 1 for r in engs])) if engs else 0) + deep_p50.append( + int(np.median([r["deepest_index"] + 1 for r in engs])) if engs else 0 + ) x = np.arange(len(labels)) - bars = ax.bar(x, eng_share, color="#dc2626", edgecolor="#7f1d1d", - linewidth=0.5, label="engaged %") + bars = ax.bar( + x, + eng_share, + color="#dc2626", + edgecolor="#7f1d1d", + linewidth=0.5, + label="engaged %", + ) for b, p, n in zip(bars, eng_share, n_calls): - ax.text(b.get_x() + b.get_width() / 2, p + 0.8, - f"{p:.0f}%\nn={n:,}", ha="center", va="bottom", fontsize=8) + ax.text( + b.get_x() + b.get_width() / 2, + p + 0.8, + f"{p:.0f}%\nn={n:,}", + ha="center", + va="bottom", + fontsize=8, + ) ax.set_xticks(x) ax.set_xticklabels(labels) ax.set_ylabel("engaged %", color="#991b1b") @@ -535,13 +640,20 @@ def plot(records: list[dict], since: str) -> Path | None: ax.set_title("engagement vs concentration (max matches in one file)") ax.set_xlabel("max matches in single file") ax2 = ax.twinx() - ax2.plot(x, deep_p50, color="#7c3aed", marker="o", linewidth=1.8, - label="p50 deepest index") + ax2.plot( + x, + deep_p50, + color="#7c3aed", + marker="o", + linewidth=1.8, + label="p50 deepest index", + ) ax2.set_ylabel("p50 deepest index", color="#5b21b6") ax2.tick_params(axis="y", labelcolor="#5b21b6") ax.grid(True, axis="y", alpha=0.25, linestyle="--") - fig.suptitle(f"search/grep result relevance — calls since {since}", - fontsize=13, y=1.0) + fig.suptitle( + f"search/grep result relevance — calls since {since}", fontsize=13, y=1.0 + ) fig.tight_layout() p = OUT_DIR / "search-relevance.png" fig.savefig(p, bbox_inches="tight") @@ -552,10 +664,14 @@ def plot(records: list[dict], since: str) -> Path | None: # --------------------------------------------------------------------------- # # Entry + def main() -> int: ap = argparse.ArgumentParser(description="search/grep result relevance analysis") - ap.add_argument("--since", default=DEFAULT_SINCE, - help=f"only calls after this date (default {DEFAULT_SINCE})") + ap.add_argument( + "--since", + default=DEFAULT_SINCE, + help=f"only calls after this date (default {DEFAULT_SINCE})", + ) args = ap.parse_args() since = datetime.strptime(args.since, "%Y-%m-%d").replace(tzinfo=timezone.utc) diff --git a/scripts/session-stats/analyze_selector_reads.py b/scripts/session-stats/analyze_selector_reads.py index d5eba1cd2..a70951859 100644 --- a/scripts/session-stats/analyze_selector_reads.py +++ b/scripts/session-stats/analyze_selector_reads.py @@ -21,6 +21,7 @@ separate cohort and excluded from interval math. Outputs: scripts/session-stats/out/selector-coverage.png """ + from __future__ import annotations import argparse @@ -50,6 +51,7 @@ _RANGE_RE = re.compile(r"^(\d+)(?:([-+])(\d+))?$") # --------------------------------------------------------------------------- # # Selector parsing + def parse_selector(path: str) -> tuple[str, int | None, int | None, str]: """Returns (base_path, start, end, kind).""" if not path: @@ -79,7 +81,9 @@ def parse_selector(path: str) -> tuple[str, int | None, int | None, str]: return base, start, start + DEFAULT_PAGE - 1, "range" -def args_to_interval(arg_json: str | None) -> tuple[str, int | None, int | None, str] | None: +def args_to_interval( + arg_json: str | None, +) -> tuple[str, int | None, int | None, str] | None: """Decode arg_json into (base, start, end, kind).""" if not arg_json: return None @@ -96,7 +100,12 @@ def args_to_interval(arg_json: str | None) -> tuple[str, int | None, int | None, # Legacy offset/limit. offset = obj.get("offset") limit = obj.get("limit") - if isinstance(offset, int) and isinstance(limit, int) and offset >= 1 and limit >= 1: + if ( + isinstance(offset, int) + and isinstance(limit, int) + and offset >= 1 + and limit >= 1 + ): return path, offset, offset + limit - 1, "range" if isinstance(offset, int) and offset >= 1: return path, offset, offset + DEFAULT_PAGE - 1, "range" @@ -108,6 +117,7 @@ def args_to_interval(arg_json: str | None) -> tuple[str, int | None, int | None, # --------------------------------------------------------------------------- # # Coverage math + def merge_intervals(ivs: list[tuple[int, int]]) -> list[tuple[int, int]]: """Merge overlapping / adjacent intervals. Inclusive bounds.""" if not ivs: @@ -123,9 +133,7 @@ def merge_intervals(ivs: list[tuple[int, int]]) -> list[tuple[int, int]]: return out -def classify_followup( - s: int, e: int, init_s: int, init_e: int -) -> str: +def classify_followup(s: int, e: int, init_s: int, init_e: int) -> str: """Where does follow-up [s,e] land relative to initial [init_s, init_e]?""" if s >= init_s and e <= init_e: return "inside" @@ -144,6 +152,7 @@ def classify_followup( # --------------------------------------------------------------------------- # # Pull + def iter_reads(conn: sqlite3.Connection, since_ms: int): sql = """ SELECT session_file, seq, timestamp, arg_json @@ -156,7 +165,9 @@ def iter_reads(conn: sqlite3.Connection, since_ms: int): def collect(conn, since_ms) -> dict[tuple[str, str], list[tuple[int, int, int, str]]]: """key (session, file) -> ordered list of (seq, start, end, kind).""" - by_key: dict[tuple[str, str], list[tuple[int, int | None, int | None, str]]] = defaultdict(list) + by_key: dict[tuple[str, str], list[tuple[int, int | None, int | None, str]]] = ( + defaultdict(list) + ) for session, seq, _ts, arg_json in iter_reads(conn, since_ms): parsed = args_to_interval(arg_json) if parsed is None: @@ -171,6 +182,7 @@ def collect(conn, since_ms) -> dict[tuple[str, str], list[tuple[int, int, int, s # --------------------------------------------------------------------------- # # Analyze + def analyze(by_key: dict) -> dict: """Compute coverage statistics over (session, file) groups whose FIRST read is a numeric range.""" @@ -213,23 +225,27 @@ def analyze(by_key: dict) -> dict: gap_lines = span - covered_lines extra_lines = max(0, covered_lines - init_size) # new lines past initial - eligible.append({ - "session": session, - "file": base, - "init_start": s0, - "init_end": e0, - "init_size": init_size, - "n_followups": len(followups), - "n_range_followups": sum(1 for k in followup_kinds if k in ("range", "default")), - "n_raw_followups": sum(1 for k in followup_kinds if k == "raw"), - "first_followup_pos": first_followup_pos, - "intervals": merged, - "regions": regions, - "covered": covered_lines, - "extra_lines": extra_lines, - "span": span, - "gap_lines": gap_lines, - }) + eligible.append( + { + "session": session, + "file": base, + "init_start": s0, + "init_end": e0, + "init_size": init_size, + "n_followups": len(followups), + "n_range_followups": sum( + 1 for k in followup_kinds if k in ("range", "default") + ), + "n_raw_followups": sum(1 for k in followup_kinds if k == "raw"), + "first_followup_pos": first_followup_pos, + "intervals": merged, + "regions": regions, + "covered": covered_lines, + "extra_lines": extra_lines, + "span": span, + "gap_lines": gap_lines, + } + ) return { "eligible": eligible, "first_kind": dict(first_kind_counts), @@ -242,18 +258,18 @@ def analyze(by_key: dict) -> dict: POS_ORDER = ["forward", "backward", "inside", "both", "gap-above", "gap-below"] POS_COLORS = { - "forward": "#2563eb", - "backward": "#0f766e", - "inside": "#9ca3af", - "both": "#7c3aed", + "forward": "#2563eb", + "backward": "#0f766e", + "inside": "#9ca3af", + "both": "#7c3aed", "gap-above": "#dc2626", "gap-below": "#d97706", } POS_HELP = { - "forward": "extended past initial end", - "backward": "extended before initial start", - "inside": "re-read inside the initial range", - "both": "extended on both sides", + "forward": "extended past initial end", + "backward": "extended before initial start", + "inside": "re-read inside the initial range", + "both": "extended on both sides", "gap-above": "disjoint hop above initial", "gap-below": "disjoint hop below initial", } @@ -269,7 +285,7 @@ def report(stats: dict) -> None: n = first_kind.get(k, 0) if n == 0: continue - print(f" {k:<10} {n:>8,} {100*n/total_pairs:>5.1f}%") + print(f" {k:<10} {n:>8,} {100 * n / total_pairs:>5.1f}%") print(f" total {total_pairs:>8,}") if not eligible: @@ -280,8 +296,12 @@ def report(stats: dict) -> None: with_followup = sum(1 for e in eligible if e["n_followups"] > 0) with_range_followup = sum(1 for e in eligible if e["n_range_followups"] > 0) print(f"\nfor {n:,} (session, file) pairs whose first read was a range:") - print(f" any follow-up read : {with_followup:>8,} ({100*with_followup/n:.1f}%)") - print(f" follow-up with a range : {with_range_followup:>8,} ({100*with_range_followup/n:.1f}%)") + print( + f" any follow-up read : {with_followup:>8,} ({100 * with_followup / n:.1f}%)" + ) + print( + f" follow-up with a range : {with_range_followup:>8,} ({100 * with_range_followup / n:.1f}%)" + ) print() print(" ----- follow-up position breakdown (all follow-up reads) -----") positions = stats["followup_pos"] @@ -290,32 +310,36 @@ def report(stats: dict) -> None: v = positions.get(k, 0) if v == 0: continue - print(f" {k:<10} {v:>8,} ({100*v/total_pos:>5.1f}%) -- {POS_HELP[k]}") + print(f" {k:<10} {v:>8,} ({100 * v / total_pos:>5.1f}%) -- {POS_HELP[k]}") # Region count distribution. regions = np.array([e["regions"] for e in eligible], dtype=np.int64) print(f"\ndisjoint regions in final coverage (per session/file):") - print(f" mean={regions.mean():.2f} median={int(np.median(regions))} " - f"p90={int(np.percentile(regions,90))} max={int(regions.max())}") + print( + f" mean={regions.mean():.2f} median={int(np.median(regions))} " + f"p90={int(np.percentile(regions, 90))} max={int(regions.max())}" + ) edges = [1, 2, 3, 4, 6, 11, 10**6] labels = ["1 (contig)", "2", "3", "4-5", "6-10", "11+"] hist, _ = np.histogram(regions, bins=edges) for label, nb in zip(labels, hist): - print(f" {label:<10} {nb:>8,} ({100*nb/regions.size:>5.1f}%)") + print(f" {label:<10} {nb:>8,} ({100 * nb / regions.size:>5.1f}%)") # Extra lines vs initial (only when follow-ups exist). fu = [e for e in eligible if e["n_range_followups"] > 0] extra = np.array([e["extra_lines"] for e in fu], dtype=np.int64) if extra.size: print(f"\nextra lines covered beyond initial range (n={extra.size:,}):") - print(f" mean={extra.mean():.0f} median={int(np.median(extra))} " - f"p75={int(np.percentile(extra,75))} p90={int(np.percentile(extra,90))} " - f"max={int(extra.max())}") + print( + f" mean={extra.mean():.0f} median={int(np.median(extra))} " + f"p75={int(np.percentile(extra, 75))} p90={int(np.percentile(extra, 90))} " + f"max={int(extra.max())}" + ) edges = [0, 1, 51, 201, 501, 2001, 10**9] labels = ["0 (no new)", "1-50", "51-200", "201-500", "501-2000", "2000+"] hist, _ = np.histogram(extra, bins=edges) for label, nb in zip(labels, hist): - print(f" {label:<12} {nb:>8,} ({100*nb/extra.size:>5.1f}%)") + print(f" {label:<12} {nb:>8,} ({100 * nb / extra.size:>5.1f}%)") # Coverage ratio. init_sizes = np.array([e["init_size"] for e in fu], dtype=np.int64) @@ -325,13 +349,16 @@ def report(stats: dict) -> None: ratio = np.where(init_sizes > 0, covered / init_sizes, np.nan) ratio = ratio[np.isfinite(ratio)] print(f"\ntotal covered / initial size:") - print(f" mean={ratio.mean():.2f}x median={np.median(ratio):.2f}x " - f"p90={np.percentile(ratio,90):.2f}x") + print( + f" mean={ratio.mean():.2f}x median={np.median(ratio):.2f}x " + f"p90={np.percentile(ratio, 90):.2f}x" + ) # --------------------------------------------------------------------------- # # Plot + def plot(stats: dict, since: str) -> Path | None: eligible = stats["eligible"] if not eligible: @@ -352,11 +379,15 @@ def plot(stats: dict, since: str) -> Path | None: colors = [POS_COLORS[k] for k in keys] bars = ax.bar(keys, vals, color=colors, edgecolor="#111", linewidth=0.5) for bar, v, k in zip(bars, vals, keys): - ax.text(bar.get_x() + bar.get_width() / 2, v + 1.0, - f"{v:.1f}%\nn={pos[k]:,}", ha="center", va="bottom", fontsize=8) - ax.set_title( - f"where do follow-up reads land vs initial range (n={total:,})" - ) + ax.text( + bar.get_x() + bar.get_width() / 2, + v + 1.0, + f"{v:.1f}%\nn={pos[k]:,}", + ha="center", + va="bottom", + fontsize=8, + ) + ax.set_title(f"where do follow-up reads land vs initial range (n={total:,})") ax.set_ylabel("share of follow-up reads") ax.set_ylim(0, max(vals) + 12) ax.yaxis.set_major_formatter(plt.FuncFormatter(lambda v, _: f"{v:.0f}%")) @@ -372,8 +403,14 @@ def plot(stats: dict, since: str) -> Path | None: colors = ["#16a34a"] + ["#2563eb"] * 5 bars = ax.bar(labels, pct, color=colors, edgecolor="#111", linewidth=0.5) for bar, p, h in zip(bars, pct, hist): - ax.text(bar.get_x() + bar.get_width() / 2, p + 1.5, - f"{p:.1f}%\nn={h:,}", ha="center", va="bottom", fontsize=8) + ax.text( + bar.get_x() + bar.get_width() / 2, + p + 1.5, + f"{p:.1f}%\nn={h:,}", + ha="center", + va="bottom", + fontsize=8, + ) ax.set_title(f"disjoint regions in final coverage (n={regions.size:,})") ax.set_ylabel("share of (session, file) pairs") ax.set_ylim(0, max(pct.max() + 12, 20)) @@ -391,8 +428,14 @@ def plot(stats: dict, since: str) -> Path | None: pct = 100 * hist / extra.size bars = ax.bar(labels, pct, color="#d97706", edgecolor="#7c2d12", linewidth=0.5) for bar, p, h in zip(bars, pct, hist): - ax.text(bar.get_x() + bar.get_width() / 2, p + 1.5, - f"{p:.1f}%\nn={h:,}", ha="center", va="bottom", fontsize=8) + ax.text( + bar.get_x() + bar.get_width() / 2, + p + 1.5, + f"{p:.1f}%\nn={h:,}", + ha="center", + va="bottom", + fontsize=8, + ) ax.set_title( f"extra lines covered beyond initial range\n" f"(only pairs that follow up, n={extra.size:,})" @@ -413,23 +456,37 @@ def plot(stats: dict, since: str) -> Path | None: ratio.sort() cdf = np.arange(1, ratio.size + 1) / ratio.size ax.plot(ratio, cdf, color="#0f766e", linewidth=1.9) - ax.axvline(1.0, color="#9ca3af", linestyle="--", linewidth=1.0, - label="covered = initial") + ax.axvline( + 1.0, + color="#9ca3af", + linestyle="--", + linewidth=1.0, + label="covered = initial", + ) for q in (0.5, 0.9): x = np.interp(q, cdf, ratio) ax.scatter([x], [q], color="#dc2626", s=22, zorder=3) - ax.annotate(f"p{int(q*100)}={x:.2f}x", (x, q), - textcoords="offset points", xytext=(6, -8), fontsize=9) + ax.annotate( + f"p{int(q * 100)}={x:.2f}x", + (x, q), + textcoords="offset points", + xytext=(6, -8), + fontsize=9, + ) ax.set_xscale("log") ax.set_xlim(0.8, max(ratio.max(), 10)) ax.set_xlabel("total covered lines / initial range (×, log)") ax.set_ylabel("CDF of pairs with follow-up") - ax.set_title(f"how much of the file does the session end up reading? (n={ratio.size:,})") + ax.set_title( + f"how much of the file does the session end up reading? (n={ratio.size:,})" + ) ax.set_ylim(0, 1.01) ax.legend(loc="lower right", frameon=False) ax.grid(True, which="both", alpha=0.25, linestyle="--") - fig.suptitle(f"selector reads — coverage map analysis (since {since})", fontsize=13, y=0.995) + fig.suptitle( + f"selector reads — coverage map analysis (since {since})", fontsize=13, y=0.995 + ) fig.tight_layout() p = OUT_DIR / "selector-coverage.png" fig.savefig(p, bbox_inches="tight") @@ -440,6 +497,7 @@ def plot(stats: dict, since: str) -> Path | None: # --------------------------------------------------------------------------- # # Examples + def dump_examples(stats: dict, k: int = 8) -> None: """Print a few coverage-map examples for sanity / intuition.""" fu = [e for e in stats["eligible"] if e["n_range_followups"] > 0] @@ -456,7 +514,11 @@ def dump_examples(stats: dict, k: int = 8) -> None: for r in sorted(buckets.keys()): candidates = buckets[r] # Prefer ones with non-default initial windows. - non_default = [c for c in candidates if (c["init_start"], c["init_end"]) != (1, DEFAULT_PAGE)] + non_default = [ + c + for c in candidates + if (c["init_start"], c["init_end"]) != (1, DEFAULT_PAGE) + ] chosen = non_default[0] if non_default else candidates[0] picks.append(chosen) if len(picks) >= k: @@ -483,20 +545,30 @@ def dump_examples(stats: dict, k: int = 8) -> None: bar[i] = "▓" bar_str = "".join(bar) file_short = e["file"][-50:] - print(f" [{bar_str}] regions={e['regions']:>2} " - f"init=[{e['init_start']},{e['init_end']}] " - f"covered={e['covered']:>4} span={e['span']:>4} {file_short}") + print( + f" [{bar_str}] regions={e['regions']:>2} " + f"init=[{e['init_start']},{e['init_end']}] " + f"covered={e['covered']:>4} span={e['span']:>4} {file_short}" + ) # --------------------------------------------------------------------------- # # Entry + def main() -> int: ap = argparse.ArgumentParser(description=__doc__.splitlines()[1]) - ap.add_argument("--since", default=DEFAULT_SINCE, - help=f"only reads at or after this date (default {DEFAULT_SINCE})") - ap.add_argument("--examples", type=int, default=8, - help="how many coverage-map examples to print (default 8)") + ap.add_argument( + "--since", + default=DEFAULT_SINCE, + help=f"only reads at or after this date (default {DEFAULT_SINCE})", + ) + ap.add_argument( + "--examples", + type=int, + default=8, + help="how many coverage-map examples to print (default 8)", + ) args = ap.parse_args() since = datetime.strptime(args.since, "%Y-%m-%d").replace(tzinfo=timezone.utc) @@ -508,8 +580,10 @@ def main() -> int: by_key = collect(conn, since_ms) conn.close() - print(f"loaded {sum(len(v) for v in by_key.values()):,} read calls across " - f"{len(by_key):,} (session, file) pairs since {args.since}") + print( + f"loaded {sum(len(v) for v in by_key.values()):,} read calls across " + f"{len(by_key):,} (session, file) pairs since {args.since}" + ) stats = analyze(by_key) report(stats) diff --git a/scripts/session-stats/harmony_backtest.py b/scripts/session-stats/harmony_backtest.py index 8b4a0c82e..14f65d1ce 100755 --- a/scripts/session-stats/harmony_backtest.py +++ b/scripts/session-stats/harmony_backtest.py @@ -25,8 +25,12 @@ DB_PATH = Path.home() / ".omp" / "stats.db" MARKER_RE = re.compile(r"\bto=functions\.[A-Za-z_][A-Za-z0-9_]*") HARMONY_RE = re.compile(r"<\|(start|end|channel|message|call|return)\|>") -CHANNEL_WORD_RE = re.compile(r"\b(analysis|commentary|assistant|user|system|developer|tool)\s+to=functions\.") -GLITCH_RE = re.compile(r"\b(changedFiles|RTLU|Jsii(?:_commentary)?|Japgolly|tRTLUfunctions|Joshi_commentary|Japgolly_commentary|jsii_commentary|Jsii_commentary|Jsii)\b") +CHANNEL_WORD_RE = re.compile( + r"\b(analysis|commentary|assistant|user|system|developer|tool)\s+to=functions\." +) +GLITCH_RE = re.compile( + r"\b(changedFiles|RTLU|Jsii(?:_commentary)?|Japgolly|tRTLUfunctions|Joshi_commentary|Japgolly_commentary|jsii_commentary|Jsii_commentary|Jsii)\b" +) NULLISH_RE = re.compile(r"\b(undefined|null)\b") BODY_CASCADE_RE = re.compile(r"\bto=functions\.[A-Za-z_][A-Za-z0-9_]*\s+code(?:\s|$)") FAKE_RESULT_RE = re.compile( @@ -38,18 +42,18 @@ FENCE_RE = re.compile(r"^\s*(```+|~~~+)") # ranges local and explicit. SCRIPT_RUN_RE = re.compile( "[" - "\u3400-\u4DBF" # CJK Extension A - "\u4E00-\u9FFF" # CJK Unified Ideographs - "\uF900-\uFAFF" # CJK Compatibility Ideographs - "\u0400-\u04FF" # Cyrillic - "\u0E00-\u0E7F" # Thai - "\u10A0-\u10FF" # Georgian - "\u0530-\u058F" # Armenian - "\u0C80-\u0CFF" # Kannada - "\u0C00-\u0C7F" # Telugu - "\u0900-\u097F" # Devanagari - "\u0600-\u06FF" # Arabic - "\u0D00-\u0D7F" # Malayalam + "\u3400-\u4dbf" # CJK Extension A + "\u4e00-\u9fff" # CJK Unified Ideographs + "\uf900-\ufaff" # CJK Compatibility Ideographs + "\u0400-\u04ff" # Cyrillic + "\u0e00-\u0e7f" # Thai + "\u10a0-\u10ff" # Georgian + "\u0530-\u058f" # Armenian + "\u0c80-\u0cff" # Kannada + "\u0c00-\u0c7f" # Telugu + "\u0900-\u097f" # Devanagari + "\u0600-\u06ff" # Arabic + "\u0d00-\u0d7f" # Malayalam "]{2,}" ) @@ -64,9 +68,15 @@ BEGIN_PATCH_RE = re.compile(r"^\*\*\* Begin Patch\s*$") END_PATCH_RE = re.compile(r"^\*\*\* End Patch\s*$") # Legacy hashline ops (kept for historical session corpus). -LEGACY_INSERT_RE = re.compile(r"^(?P[«»])\s*(?PBOF|EOF|[1-9][0-9]*[A-Za-z]{2})\s*$") -LEGACY_RANGE_RE = re.compile(r"(?P[1-9][0-9]*[A-Za-z]{2})(?:\.\.(?P[1-9][0-9]*[A-Za-z]{2}))?") -LEGACY_REPLACE_RE = re.compile(r"^≔\s*(?P[1-9][0-9]*[A-Za-z]{2}(?:\.\.[1-9][0-9]*[A-Za-z]{2})?)\s*$") +LEGACY_INSERT_RE = re.compile( + r"^(?P[«»])\s*(?PBOF|EOF|[1-9][0-9]*[A-Za-z]{2})\s*$" +) +LEGACY_RANGE_RE = re.compile( + r"(?P[1-9][0-9]*[A-Za-z]{2})(?:\.\.(?P[1-9][0-9]*[A-Za-z]{2}))?" +) +LEGACY_REPLACE_RE = re.compile( + r"^≔\s*(?P[1-9][0-9]*[A-Za-z]{2}(?:\.\.[1-9][0-9]*[A-Za-z]{2})?)\s*$" +) # Current hashline ops. NEW_INSERT_RE = re.compile( @@ -183,7 +193,6 @@ def commas(n: int) -> str: return f"{n:,}" - def one_line(text: str, limit: int = 180) -> str: text = text.replace("\r", "\\r").replace("\n", " | ").replace("\t", "\\t") if len(text) <= limit: @@ -242,7 +251,9 @@ def marker_evidence_for( window16 = text[max(0, start - 16) : min(len(text), end + 16)] window200 = text[start : min(len(text), start + 200)] - for c in CHANNEL_WORD_RE.finditer(text[max(0, start - 64) : min(len(text), end + 16)]): + for c in CHANNEL_WORD_RE.finditer( + text[max(0, start - 64) : min(len(text), end + 16)] + ): absolute_start = max(0, start - 64) + c.start() absolute_end = max(0, start - 64) + c.end() if absolute_start <= start < absolute_end: @@ -255,7 +266,9 @@ def marker_evidence_for( ev.classes.add("N") if script_mismatch_near(text, start, end): ev.classes.add("S") - if BODY_CASCADE_RE.match(window200) and MARKER_RE.search(window200[marker.end() - start :]): + if BODY_CASCADE_RE.match(window200) and MARKER_RE.search( + window200[marker.end() - start :] + ): ev.classes.add("B") if FAKE_RESULT_RE.match(text, start): ev.classes.add("R") @@ -280,7 +293,9 @@ def detect_signals( signals.append(Signal("H", h.start(), h.end(), h.group(0))) for marker in MARKER_RE.finditer(text): - ev = marker_evidence_for(text, marker, parsed_end, respect_fences, include_nullish) + ev = marker_evidence_for( + text, marker, parsed_end, respect_fences, include_nullish + ) if ev is None: continue marker_evidence.append(ev) @@ -319,6 +334,7 @@ def line_spans(text: str) -> list[tuple[str, int, int]]: # splitlines(keepends=True) already includes the final unterminated line. return out + def parse_legacy_diff_boundary(text: str, *, loose_tail: bool = False) -> EditBoundary: """Best-effort parser for pre-hashline edit inputs. @@ -342,7 +358,12 @@ def parse_legacy_diff_boundary(text: str, *, loose_tail: bool = False) -> EditBo for line, _start, end in line_spans(text): line_no += 1 - if loose_tail and (MARKER_RE.search(line) or HARMONY_RE.search(line)) and cur is not None and cur.op_count > 0: + if ( + loose_tail + and (MARKER_RE.search(line) or HARMONY_RE.search(line)) + and cur is not None + and cur.op_count > 0 + ): break if line.startswith("---"): @@ -395,7 +416,9 @@ def parse_legacy_diff_boundary(text: str, *, loose_tail: bool = False) -> EditBo in_payload = True continue - if in_payload and (line.startswith("-") or line.startswith(" ") or line.startswith("\\")): + if in_payload and ( + line.startswith("-") or line.startswith(" ") or line.startswith("\\") + ): if line.startswith("-") and not line.startswith("---"): cur.deleted_lines += 1 parsed_end = end @@ -415,7 +438,9 @@ def parse_legacy_diff_boundary(text: str, *, loose_tail: bool = False) -> EditBo return EditBoundary( ok=parsed_end > 0 and bool(sections), parsed_end=parsed_end, - reason="legacy-edit-ok" if parsed_end > 0 and sections else "no-complete-edit-prefix", + reason="legacy-edit-ok" + if parsed_end > 0 and sections + else "no-complete-edit-prefix", sections=sections, line_no=line_no, ) @@ -547,7 +572,9 @@ def parse_edit_boundary(text: str, *, legacy_loose_tail: bool = False) -> EditBo if rng: sigil = rng.group("sigil") cur.op_count += 1 - cur.deleted_lines += new_range_deleted_lines(rng.group("a"), rng.group("b")) + cur.deleted_lines += new_range_deleted_lines( + rng.group("a"), rng.group("b") + ) if sigil == ":" and rng.group("inline"): cur.payload_lines += 1 parsed_end = end @@ -620,19 +647,36 @@ def parse_arg_json(raw: str) -> tuple[Any | None, bool, str]: return None, False, f"json-error:{exc.pos}:{exc.msg}" -def extract_primary_text(tool_name: str, arg_json: str, parsed: Any | None) -> tuple[str, str]: - if tool_name == "edit" and isinstance(parsed, dict) and isinstance(parsed.get("input"), str): +def extract_primary_text( + tool_name: str, arg_json: str, parsed: Any | None +) -> tuple[str, str]: + if ( + tool_name == "edit" + and isinstance(parsed, dict) + and isinstance(parsed.get("input"), str) + ): return "edit.input", parsed["input"] - if tool_name == "eval" and isinstance(parsed, dict) and isinstance(parsed.get("input"), str): + if ( + tool_name == "eval" + and isinstance(parsed, dict) + and isinstance(parsed.get("input"), str) + ): return "eval.input", parsed["input"] - if tool_name == "write" and isinstance(parsed, dict) and isinstance(parsed.get("content"), str): + if ( + tool_name == "write" + and isinstance(parsed, dict) + and isinstance(parsed.get("content"), str) + ): return "write.content", parsed["content"] - if tool_name == "bash" and isinstance(parsed, dict) and isinstance(parsed.get("command"), str): + if ( + tool_name == "bash" + and isinstance(parsed, dict) + and isinstance(parsed.get("command"), str) + ): return "bash.command", parsed["command"] return "arg_json", arg_json - def action_for_tool( tool_name: str, surface: str, @@ -642,7 +686,12 @@ def action_for_tool( ) -> str: if not signals: return "allow" - if tool_name == "edit" and surface == "edit.input" and boundary is not None and boundary.ok: + if ( + tool_name == "edit" + and surface == "edit.input" + and boundary is not None + and boundary.ok + ): if all(sig.start >= boundary.parsed_end for sig in signals): return "sanitize_tail" return "abort_replay" @@ -782,7 +831,9 @@ def candidate_where(column: str) -> str: ) -def scan_tools(conn: sqlite3.Connection, args: argparse.Namespace) -> list[ToolBacktest]: +def scan_tools( + conn: sqlite3.Connection, args: argparse.Namespace +) -> list[ToolBacktest]: where = candidate_where("arg_json") params: list[Any] = [] if args.provider: @@ -814,7 +865,9 @@ def scan_tools(conn: sqlite3.Connection, args: argparse.Namespace) -> list[ToolB ] -def scan_assistant(conn: sqlite3.Connection, args: argparse.Namespace) -> list[TextBacktest]: +def scan_assistant( + conn: sqlite3.Connection, args: argparse.Namespace +) -> list[TextBacktest]: if not args.include_assistant: return [] text_where = candidate_where("text_blob") @@ -871,8 +924,13 @@ def print_tool_summary(results: list[ToolBacktest]) -> None: print("=== tool-call scan ===") print(f"candidate rows: {commas(len(results))}") print_counter("\nby action:", Counter(r.action for r in results)) - print_counter("\nby tool/action:", Counter(f"{r.tool_name}:{r.action}" for r in results)) - print_counter("\nby model/action:", Counter(f"{r.model or ''}:{r.action}" for r in results)) + print_counter( + "\nby tool/action:", Counter(f"{r.tool_name}:{r.action}" for r in results) + ) + print_counter( + "\nby model/action:", + Counter(f"{r.model or ''}:{r.action}" for r in results), + ) signal_counter: Counter[str] = Counter() for r in results: if r.signals: @@ -890,9 +948,17 @@ def print_tool_summary(results: list[ToolBacktest]) -> None: print(f" sanitize_tail: {commas(sanitized)}") print(f" abort_replay: {commas(aborted)}") if sanitized: - preserved_ops = sum(r.edit_ops for r in edit_results if r.action == "sanitize_tail") - preserved_payload = sum(r.edit_payload_lines for r in edit_results if r.action == "sanitize_tail") - removed = sum(r.removed_len for r in edit_results if r.action == "sanitize_tail") + preserved_ops = sum( + r.edit_ops for r in edit_results if r.action == "sanitize_tail" + ) + preserved_payload = sum( + r.edit_payload_lines + for r in edit_results + if r.action == "sanitize_tail" + ) + removed = sum( + r.removed_len for r in edit_results if r.action == "sanitize_tail" + ) print(f" ops preserved by sanitize: {commas(preserved_ops)}") print(f" payload lines preserved: {commas(preserved_payload)}") print(f" tail bytes removed: {commas(removed)}") @@ -904,28 +970,40 @@ def print_text_summary(results: list[TextBacktest]) -> None: print("\n=== assistant message scan ===") print(f"candidate surfaces: {commas(len(results))}") print_counter("\nby action:", Counter(r.action for r in results)) - print_counter("\nby surface/action:", Counter(f"{r.surface}:{r.action}" for r in results)) - print_counter("\nby model/action:", Counter(f"{r.model or ''}:{r.action}" for r in results)) - + print_counter( + "\nby surface/action:", Counter(f"{r.surface}:{r.action}" for r in results) + ) + print_counter( + "\nby model/action:", + Counter(f"{r.model or ''}:{r.action}" for r in results), + ) def signal_summary(labels: list[str], limit: int = 6) -> str: if not labels: return "none" counts = Counter(labels) - parts = [f"{label}x{count}" if count > 1 else label for label, count in counts.most_common(limit)] + parts = [ + f"{label}x{count}" if count > 1 else label + for label, count in counts.most_common(limit) + ] rest = sum(counts.values()) - sum(count for _, count in counts.most_common(limit)) if rest: parts.append(f"+{rest} more") return ",".join(parts) + def print_examples(results: list[ToolBacktest], show: int) -> None: if show <= 0: return print(f"\n=== sanitize_tail edit examples (up to {show}) ===") - sanitize_examples = [r for r in results if r.tool_name == "edit" and r.action == "sanitize_tail"] + sanitize_examples = [ + r for r in results if r.tool_name == "edit" and r.action == "sanitize_tail" + ] for r in sanitize_examples[:show]: - print(f"\n[id={r.row_id} seq={r.seq} model={r.model} signals={signal_summary(r.signals)}]") + print( + f"\n[id={r.row_id} seq={r.seq} model={r.model} signals={signal_summary(r.signals)}]" + ) print(f"session: {r.session_file}") print(f"file(s): {', '.join(r.edit_files) if r.edit_files else ''}") print( @@ -938,7 +1016,9 @@ def print_examples(results: list[ToolBacktest], show: int) -> None: print(f"\n=== abort_replay examples (up to {show}) ===") abort_examples = [r for r in results if r.action == "abort_replay"] for r in abort_examples[:show]: - print(f"\n[id={r.row_id} tool={r.tool_name} surface={r.surface} seq={r.seq} model={r.model} signals={signal_summary(r.signals)}]") + print( + f"\n[id={r.row_id} tool={r.tool_name} surface={r.surface} seq={r.seq} model={r.model} signals={signal_summary(r.signals)}]" + ) print(f"session: {r.session_file}") if r.tool_name == "edit": print( @@ -948,7 +1028,9 @@ def print_examples(results: list[ToolBacktest], show: int) -> None: print(f"context: {r.context_preview}") -def write_json_report(path: Path, tools: list[ToolBacktest], texts: list[TextBacktest]) -> None: +def write_json_report( + path: Path, tools: list[ToolBacktest], texts: list[TextBacktest] +) -> None: def tool_dict(r: ToolBacktest) -> dict[str, Any]: return { "surface": r.surface, @@ -998,7 +1080,9 @@ def write_json_report(path: Path, tools: list[ToolBacktest], texts: list[TextBac "assistant_surfaces": [text_dict(r) for r in texts], } path.parent.mkdir(parents=True, exist_ok=True) - path.write_text(json.dumps(payload, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") + path.write_text( + json.dumps(payload, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" + ) def main() -> int: @@ -1016,22 +1100,46 @@ def main() -> int: "tail = only H or marker after parsed boundary" ), ) - ap.add_argument("--provider", default=None, help="restrict tool/message rows to a provider") - ap.add_argument("--model", default=None, help="restrict tool/message rows to a model") + ap.add_argument( + "--provider", default=None, help="restrict tool/message rows to a provider" + ) + ap.add_argument( + "--model", default=None, help="restrict tool/message rows to a model" + ) ap.add_argument("--tool", default=None, help="restrict tool-call rows to one tool") - ap.add_argument("--include-assistant", action="store_true", help="also scan assistant text/thinking surfaces") - ap.add_argument("--include-nullish", action="store_true", help="treat adjacent null/undefined as signal N") - ap.add_argument("--no-fence-context", action="store_true", help="do not exempt Markdown fenced blocks") - ap.add_argument("--legacy-loose-tail", action="store_true", help="model old raw-payload edit inputs as tail-sanitizable at first marker line") + ap.add_argument( + "--include-assistant", + action="store_true", + help="also scan assistant text/thinking surfaces", + ) + ap.add_argument( + "--include-nullish", + action="store_true", + help="treat adjacent null/undefined as signal N", + ) + ap.add_argument( + "--no-fence-context", + action="store_true", + help="do not exempt Markdown fenced blocks", + ) + ap.add_argument( + "--legacy-loose-tail", + action="store_true", + help="model old raw-payload edit inputs as tail-sanitizable at first marker line", + ) ap.add_argument("--show", type=int, default=8, help="examples per action group") - ap.add_argument("--json-out", type=Path, default=None, help="write machine-readable report") + ap.add_argument( + "--json-out", type=Path, default=None, help="write machine-readable report" + ) args = ap.parse_args() conn = open_ro(args.db) print("=== harmony leak backtest ===") print(f"db: {args.db}") print(f"strategy: {args.strategy}") - print(f"fences: {'ignored for action' if not args.no_fence_context else 'scanned as active text'}") + print( + f"fences: {'ignored for action' if not args.no_fence_context else 'scanned as active text'}" + ) if args.legacy_loose_tail: print("legacy: loose tail mode") if args.provider: diff --git a/scripts/session-stats/optimize_read_config.py b/scripts/session-stats/optimize_read_config.py index bccf5d029..b9d7dba1a 100644 --- a/scripts/session-stats/optimize_read_config.py +++ b/scripts/session-stats/optimize_read_config.py @@ -23,6 +23,7 @@ Output: scripts/session-stats/out/read-config-sweep.png console table with the recommended config + savings """ + from __future__ import annotations import argparse @@ -64,19 +65,18 @@ ROUNDTRIP_OVERHEAD = 200 # Selector parser (reuses the same rules as analyze_selector_reads.py). _RANGE_RE = re.compile(r"^(\d+)(?:([-+])(\d+))?$") -_FOOTER_RE = re.compile( - r"\[Showing lines (\d+)-(\d+) of (\d+)\." -) +_FOOTER_RE = re.compile(r"\[Showing lines (\d+)-(\d+) of (\d+)\.") _TRUNCATED_RE = re.compile(r"\[Output truncated") # --------------------------------------------------------------------------- # # Selector → intent + class Intent(NamedTuple): - kind: str # 'bare' | 'range' | 'raw' | 'conflicts' | 'other' - start: int | None # requested start line (1-indexed) — only meaningful for 'range' - end: int | None # requested end line (1-indexed, inclusive) — None = open-ended + kind: str # 'bare' | 'range' | 'raw' | 'conflicts' | 'other' + start: int | None # requested start line (1-indexed) — only meaningful for 'range' + end: int | None # requested end line (1-indexed, inclusive) — None = open-ended def parse_selector(path: str) -> tuple[str, Intent]: @@ -119,7 +119,12 @@ def parse_args(arg_json: str | None) -> tuple[str | None, Intent]: # Legacy offset/limit treated as an explicit range. offset = obj.get("offset") limit = obj.get("limit") - if isinstance(offset, int) and isinstance(limit, int) and offset >= 1 and limit >= 1: + if ( + isinstance(offset, int) + and isinstance(limit, int) + and offset >= 1 + and limit >= 1 + ): return path, Intent("range", offset, offset + limit - 1) if isinstance(offset, int) and offset >= 1: return path, Intent("range", offset, None) @@ -129,6 +134,7 @@ def parse_args(arg_json: str | None) -> tuple[str | None, Intent]: # --------------------------------------------------------------------------- # # Footer parser → (returned_start, returned_end, file_total_lines) + def parse_footer(tail: str | None) -> tuple[int | None, int | None, int | None, bool]: """Returns (returned_a, returned_b, file_total, was_byte_truncated).""" if not tail: @@ -136,13 +142,18 @@ def parse_footer(tail: str | None) -> tuple[int | None, int | None, int | None, m = _FOOTER_RE.search(tail) if not m: return None, None, None, bool(_TRUNCATED_RE.search(tail)) - return (int(m.group(1)), int(m.group(2)), int(m.group(3)), - bool(_TRUNCATED_RE.search(tail))) + return ( + int(m.group(1)), + int(m.group(2)), + int(m.group(3)), + bool(_TRUNCATED_RE.search(tail)), + ) # --------------------------------------------------------------------------- # # Coverage utilities + def merge(ivs: list[tuple[int, int]]) -> list[tuple[int, int]]: if not ivs: return [] @@ -186,15 +197,16 @@ def subtract(s: int, e: int, ivs: list[tuple[int, int]]) -> list[tuple[int, int] # --------------------------------------------------------------------------- # # Data model + class ReadCall(NamedTuple): seq: int intent: Intent base: str - actual_a: int | None # what came back: lines [actual_a, actual_b] + actual_a: int | None # what came back: lines [actual_a, actual_b] actual_b: int | None - file_total: int | None # from footer - tokens: int # observed result tokens - was_truncated: bool # [Output truncated marker present + file_total: int | None # from footer + tokens: int # observed result tokens + was_truncated: bool # [Output truncated marker present def fetch_reads(conn: sqlite3.Connection, since_ms: int) -> list[tuple[str, ReadCall]]: @@ -223,16 +235,30 @@ def fetch_reads(conn: sqlite3.Connection, since_ms: int) -> list[tuple[str, Read if not base or base.endswith("/") or "://" in base: continue actual_a, actual_b, file_total, was_trunc = parse_footer(tail) - out.append((session, ReadCall(seq, intent, base, actual_a, actual_b, - file_total, int(tokens), was_trunc))) + out.append( + ( + session, + ReadCall( + seq, + intent, + base, + actual_a, + actual_b, + file_total, + int(tokens), + was_trunc, + ), + ) + ) return out # --------------------------------------------------------------------------- # # Per-file aggregates + class FileStats(NamedTuple): - size_lines: int # best-effort estimate + size_lines: int # best-effort estimate tokens_per_line: float bytes_per_line: float # only when we can derive (currently we can't, so fallback) @@ -271,11 +297,7 @@ def aggregate_files(reads: list[tuple[str, ReadCall]]) -> dict[str, FileStats]: by_file_tok_lines[rc.base].append((rc.tokens, n)) out: dict[str, FileStats] = {} - files = ( - set(by_file_total_lines) - | set(by_file_max_end) - | set(by_file_tok_lines) - ) + files = set(by_file_total_lines) | set(by_file_max_end) | set(by_file_tok_lines) for f in files: size = by_file_total_lines.get(f) or by_file_max_end.get(f, 1) tok_lines = by_file_tok_lines.get(f, []) @@ -285,15 +307,19 @@ def aggregate_files(reads: list[tuple[str, ReadCall]]) -> dict[str, FileStats]: tpl = tot_tok / max(tot_ln, 1) else: tpl = FALLBACK_TPL - out[f] = FileStats(size_lines=size, tokens_per_line=tpl, - bytes_per_line=max(8.0, tpl * 4.0)) + out[f] = FileStats( + size_lines=size, tokens_per_line=tpl, bytes_per_line=max(8.0, tpl * 4.0) + ) return out # --------------------------------------------------------------------------- # # Per-pair grouping -def group_pairs(reads: list[tuple[str, ReadCall]]) -> dict[tuple[str, str], list[ReadCall]]: + +def group_pairs( + reads: list[tuple[str, ReadCall]], +) -> dict[tuple[str, str], list[ReadCall]]: by_pair: dict[tuple[str, str], list[ReadCall]] = defaultdict(list) for session, rc in reads: by_pair[(session, rc.base)].append(rc) @@ -304,15 +330,18 @@ def group_pairs(reads: list[tuple[str, ReadCall]]) -> dict[tuple[str, str], list # --------------------------------------------------------------------------- # # Simulator + class Config(NamedTuple): - default_page: int # lines returned for a bare path read - line_cap: int # absolute max lines per read - byte_cap: int # max bytes per read (modelled as line cap via bytes_per_line) + default_page: int # lines returned for a bare path read + line_cap: int # absolute max lines per read + byte_cap: int # max bytes per read (modelled as line cap via bytes_per_line) summarize_min: int # min file size (lines) for summarizer to fire on bare reads - # (-1 disables summarizer; 0 = always) + # (-1 disables summarizer; 0 = always) -def effective_returned(rc: ReadCall, fs: FileStats, cfg: Config) -> tuple[int, int] | None: +def effective_returned( + rc: ReadCall, fs: FileStats, cfg: Config +) -> tuple[int, int] | None: """Range the tool actually returns for one call under cfg. Honours intent (what the agent asked for), then applies (default_page, @@ -330,7 +359,9 @@ def effective_returned(rc: ReadCall, fs: FileStats, cfg: Config) -> tuple[int, i start, end_intent = 1, cfg.default_page elif intent.kind == "range": start = intent.start or 1 - end_intent = intent.end if intent.end is not None else (start + cfg.default_page - 1) + end_intent = ( + intent.end if intent.end is not None else (start + cfg.default_page - 1) + ) elif intent.kind == "raw": start, end_intent = 1, size else: @@ -343,11 +374,17 @@ def effective_returned(rc: ReadCall, fs: FileStats, cfg: Config) -> tuple[int, i return (start, end) -def cost_of_chunk(start: int, end: int, fs: FileStats, intent_kind: str, cfg: Config) -> float: +def cost_of_chunk( + start: int, end: int, fs: FileStats, intent_kind: str, cfg: Config +) -> float: """Estimated result tokens for returning [start, end] of this file.""" span = max(end - start + 1, 0) raw = span * fs.tokens_per_line - if intent_kind == "bare" and cfg.summarize_min >= 0 and fs.size_lines >= cfg.summarize_min: + if ( + intent_kind == "bare" + and cfg.summarize_min >= 0 + and fs.size_lines >= cfg.summarize_min + ): # Calibrated from observed post-deploy summary-eligible reads: # tokens/line collapses to ~0.35× the verbatim rate. return raw * 0.35 @@ -387,7 +424,11 @@ def replay_pair(reads: list[ReadCall], fs: FileStats, cfg: Config) -> tuple[floa total = 0.0 kept = 0 for rc in reads: - if rc.actual_a is not None and rc.actual_b is not None and rc.actual_b >= rc.actual_a: + if ( + rc.actual_a is not None + and rc.actual_b is not None + and rc.actual_b >= rc.actual_a + ): observed_needed.append((rc.actual_a, rc.actual_b)) ret = effective_returned(rc, fs, cfg) if ret is None: @@ -415,12 +456,15 @@ def replay_pair(reads: list[ReadCall], fs: FileStats, cfg: Config) -> tuple[floa return total, kept -def simulate(by_pair: dict, files: dict[str, FileStats], cfg: Config) -> tuple[float, int]: +def simulate( + by_pair: dict, files: dict[str, FileStats], cfg: Config +) -> tuple[float, int]: grand = 0.0 kept = 0 for (_session, base), reads in by_pair.items(): - fs = files.get(base) or FileStats(size_lines=1, tokens_per_line=FALLBACK_TPL, - bytes_per_line=FALLBACK_BPL) + fs = files.get(base) or FileStats( + size_lines=1, tokens_per_line=FALLBACK_TPL, bytes_per_line=FALLBACK_BPL + ) t, k = replay_pair(reads, fs, cfg) grand += t kept += k @@ -436,6 +480,7 @@ def baseline_observed(reads: list[tuple[str, ReadCall]]) -> tuple[int, int]: # --------------------------------------------------------------------------- # # Sweep + report + def sweep(by_pair: dict, files: dict[str, FileStats]) -> dict: defaults = [200, 300, 400, 500, 700, 1000, 1500, 2000, 3000] line_caps = [500, 1000, 1500, 2000, 3000, 5000] @@ -445,8 +490,9 @@ def sweep(by_pair: dict, files: dict[str, FileStats]) -> dict: grid_calls = np.zeros((len(defaults), len(line_caps)), dtype=np.int64) for i, D in enumerate(defaults): for j, L in enumerate(line_caps): - cfg = Config(default_page=D, line_cap=L, - byte_cap=CURRENT_BYTE_CAP, summarize_min=0) + cfg = Config( + default_page=D, line_cap=L, byte_cap=CURRENT_BYTE_CAP, summarize_min=0 + ) t, k = simulate(by_pair, files, cfg) grid_tokens[i, j] = t grid_calls[i, j] = k @@ -459,27 +505,46 @@ def sweep(by_pair: dict, files: dict[str, FileStats]) -> dict: # Sweep summarize_min at best (D, L). sm_tokens = [] for sm in summary_thresholds: - cfg = Config(default_page=best_DL[0], line_cap=best_DL[1], - byte_cap=CURRENT_BYTE_CAP, summarize_min=sm) + cfg = Config( + default_page=best_DL[0], + line_cap=best_DL[1], + byte_cap=CURRENT_BYTE_CAP, + summarize_min=sm, + ) t, k = simulate(by_pair, files, cfg) sm_tokens.append((sm, t, k)) best_sm = min(sm_tokens, key=lambda x: x[1]) # Sweep byte_cap at best (D, L, summarize_min). - byte_caps = [16 * 1024, 32 * 1024, 50 * 1024, 75 * 1024, 100 * 1024, - 150 * 1024, 200 * 1024] + byte_caps = [ + 16 * 1024, + 32 * 1024, + 50 * 1024, + 75 * 1024, + 100 * 1024, + 150 * 1024, + 200 * 1024, + ] bc_tokens = [] for bc in byte_caps: - cfg = Config(default_page=best_DL[0], line_cap=best_DL[1], - byte_cap=bc, summarize_min=best_sm[0]) + cfg = Config( + default_page=best_DL[0], + line_cap=best_DL[1], + byte_cap=bc, + summarize_min=best_sm[0], + ) t, k = simulate(by_pair, files, cfg) bc_tokens.append((bc, t, k)) best_bc = min(bc_tokens, key=lambda x: x[1]) # Final combined config (D, L, summarize_min, byte_cap) — should be the # global minimum given the order of dimensions. - final_cfg = Config(default_page=best_DL[0], line_cap=best_DL[1], - byte_cap=best_bc[0], summarize_min=best_sm[0]) + final_cfg = Config( + default_page=best_DL[0], + line_cap=best_DL[1], + byte_cap=best_bc[0], + summarize_min=best_sm[0], + ) final_tokens, final_calls = simulate(by_pair, files, final_cfg) return { @@ -501,6 +566,7 @@ def sweep(by_pair: dict, files: dict[str, FileStats]) -> dict: # --------------------------------------------------------------------------- # # Plotting + def plot(result: dict, baseline_sim: float, observed: int, out_path: Path) -> None: OUT_DIR.mkdir(parents=True, exist_ok=True) plt.rcParams.update({"figure.dpi": 110, "font.size": 10}) @@ -511,8 +577,14 @@ def plot(result: dict, baseline_sim: float, observed: int, out_path: Path) -> No ax = axes[0, 0] grid = result["grid_tokens"] rel = grid / baseline_sim - im = ax.imshow(rel, cmap="RdYlGn_r", aspect="auto", origin="lower", - vmin=max(0.6, rel.min()), vmax=min(1.4, rel.max() + 0.02)) + im = ax.imshow( + rel, + cmap="RdYlGn_r", + aspect="auto", + origin="lower", + vmin=max(0.6, rel.min()), + vmax=min(1.4, rel.max() + 0.02), + ) ax.set_xticks(range(len(result["line_caps"]))) ax.set_xticklabels(result["line_caps"]) ax.set_yticks(range(len(result["defaults"]))) @@ -522,21 +594,53 @@ def plot(result: dict, baseline_sim: float, observed: int, out_path: Path) -> No ax.set_title("simulated read tokens / baseline\n(green = cheaper, red = more)") for i in range(grid.shape[0]): for j in range(grid.shape[1]): - ax.text(j, i, f"{rel[i,j]:.2f}", ha="center", va="center", - color="black", fontsize=8) + ax.text( + j, + i, + f"{rel[i, j]:.2f}", + ha="center", + va="center", + color="black", + fontsize=8, + ) fig.colorbar(im, ax=ax, fraction=0.05) # Highlight current and best. - cur_i = result["defaults"].index(CURRENT_DEFAULT) if CURRENT_DEFAULT in result["defaults"] else None - cur_j = result["line_caps"].index(CURRENT_LINE_CAP) if CURRENT_LINE_CAP in result["line_caps"] else None + cur_i = ( + result["defaults"].index(CURRENT_DEFAULT) + if CURRENT_DEFAULT in result["defaults"] + else None + ) + cur_j = ( + result["line_caps"].index(CURRENT_LINE_CAP) + if CURRENT_LINE_CAP in result["line_caps"] + else None + ) if cur_i is not None and cur_j is not None: - ax.add_patch(mpatches.Rectangle((cur_j - 0.5, cur_i - 0.5), 1, 1, - fill=False, edgecolor="#1d4ed8", - linewidth=2.4, label="current")) + ax.add_patch( + mpatches.Rectangle( + (cur_j - 0.5, cur_i - 0.5), + 1, + 1, + fill=False, + edgecolor="#1d4ed8", + linewidth=2.4, + label="current", + ) + ) best_i = result["defaults"].index(result["best_DL"][0]) best_j = result["line_caps"].index(result["best_DL"][1]) - ax.add_patch(mpatches.Rectangle((best_j - 0.5, best_i - 0.5), 1, 1, - fill=False, edgecolor="#000", - linewidth=2.4, linestyle="--", label="optimum")) + ax.add_patch( + mpatches.Rectangle( + (best_j - 0.5, best_i - 0.5), + 1, + 1, + fill=False, + edgecolor="#000", + linewidth=2.4, + linestyle="--", + label="optimum", + ) + ) ax.legend(loc="upper right", frameon=True, fontsize=9) # Default-page line (at best line cap). @@ -546,9 +650,17 @@ def plot(result: dict, baseline_sim: float, observed: int, out_path: Path) -> No col = grid[:, j] / baseline_sim ax.plot(result["defaults"], col, marker="o", linewidth=1.8, color="#0f766e") ax.axhline(1.0, color="#9ca3af", linestyle="--", linewidth=1.0) - ax.axvline(CURRENT_DEFAULT, color="#1d4ed8", linestyle=":", linewidth=1.2, label="current default") + ax.axvline( + CURRENT_DEFAULT, + color="#1d4ed8", + linestyle=":", + linewidth=1.2, + label="current default", + ) best_D = result["best_DL"][0] - ax.axvline(best_D, color="#000", linestyle="--", linewidth=1.4, label=f"optimum D={best_D}") + ax.axvline( + best_D, color="#000", linestyle="--", linewidth=1.4, label=f"optimum D={best_D}" + ) ax.set_xscale("log") ax.set_xlabel("default page (D) — log scale") ax.set_ylabel("simulated tokens / baseline") @@ -559,17 +671,33 @@ def plot(result: dict, baseline_sim: float, observed: int, out_path: Path) -> No # Summarizer threshold sweep. ax = axes[1, 0] sm_data = result["summary_sweep"] - xs = [str("off") if sm == -1 else ("always" if sm == 0 else f"≥{sm}") for sm, _, _ in sm_data] + xs = [ + str("off") if sm == -1 else ("always" if sm == 0 else f"≥{sm}") + for sm, _, _ in sm_data + ] ys = [t / baseline_sim for _, t, _ in sm_data] - bars = ax.bar(xs, ys, color=["#dc2626" if y > 1 else "#0f766e" for y in ys], - edgecolor="#111", linewidth=0.5) + bars = ax.bar( + xs, + ys, + color=["#dc2626" if y > 1 else "#0f766e" for y in ys], + edgecolor="#111", + linewidth=0.5, + ) for bar, y in zip(bars, ys): - ax.text(bar.get_x() + bar.get_width() / 2, y + 0.005, - f"{(y - 1) * 100:+.1f}%", ha="center", va="bottom", fontsize=9) + ax.text( + bar.get_x() + bar.get_width() / 2, + y + 0.005, + f"{(y - 1) * 100:+.1f}%", + ha="center", + va="bottom", + fontsize=9, + ) ax.axhline(1.0, color="#9ca3af", linestyle="--", linewidth=1.0) ax.set_ylabel("simulated tokens / baseline") ax.set_xlabel("summarize files ≥ N lines") - ax.set_title(f"summarizer threshold sweep (D={result['best_DL'][0]}, L={result['best_DL'][1]})") + ax.set_title( + f"summarizer threshold sweep (D={result['best_DL'][0]}, L={result['best_DL'][1]})" + ) ax.grid(True, axis="y", alpha=0.25, linestyle="--") # Byte cap sweep. @@ -579,11 +707,21 @@ def plot(result: dict, baseline_sim: float, observed: int, out_path: Path) -> No ys = [t / baseline_sim for _, t, _ in bc_data] ax.plot(xs_kb, ys, marker="o", linewidth=1.8, color="#7c3aed") ax.axhline(1.0, color="#9ca3af", linestyle="--", linewidth=1.0) - ax.axvline(CURRENT_BYTE_CAP / 1024, color="#1d4ed8", linestyle=":", - linewidth=1.2, label="current byte cap") + ax.axvline( + CURRENT_BYTE_CAP / 1024, + color="#1d4ed8", + linestyle=":", + linewidth=1.2, + label="current byte cap", + ) best_bc_kb = result["best_byte_cap"][0] // 1024 - ax.axvline(best_bc_kb, color="#000", linestyle="--", linewidth=1.4, - label=f"optimum {best_bc_kb} KB") + ax.axvline( + best_bc_kb, + color="#000", + linestyle="--", + linewidth=1.4, + label=f"optimum {best_bc_kb} KB", + ) ax.set_xlabel("byte cap (KB)") ax.set_ylabel("simulated tokens / baseline") ax.set_title("sensitivity to byte cap") @@ -593,7 +731,8 @@ def plot(result: dict, baseline_sim: float, observed: int, out_path: Path) -> No fig.suptitle( f"read tool config sweep — observed read spend {observed:,}, " f"simulator baseline {baseline_sim:,.0f}", - fontsize=12, y=1.02, + fontsize=12, + y=1.02, ) fig.tight_layout() fig.savefig(out_path, bbox_inches="tight") @@ -603,14 +742,20 @@ def plot(result: dict, baseline_sim: float, observed: int, out_path: Path) -> No # --------------------------------------------------------------------------- # # Report + def fmt_pct(x: float) -> str: if x >= 0: - return f"+{x*100:.1f}%" - return f"{x*100:.1f}%" + return f"+{x * 100:.1f}%" + return f"{x * 100:.1f}%" -def report(result: dict, baseline_sim: float, baseline_calls: int, - observed: int, observed_calls: int) -> None: +def report( + result: dict, + baseline_sim: float, + baseline_calls: int, + observed: int, + observed_calls: int, +) -> None: defaults = result["defaults"] line_caps = result["line_caps"] grid = result["grid_tokens"] @@ -618,8 +763,10 @@ def report(result: dict, baseline_sim: float, baseline_calls: int, print(f"\nbaseline (current config: D={CURRENT_DEFAULT}, L={CURRENT_LINE_CAP}):") print(f" observed result tokens = {observed:>13,} (truth)") - print(f" simulator under baseline = {baseline_sim:>13,.0f} " - f"({fmt_pct((baseline_sim - observed) / observed)} vs observed)") + print( + f" simulator under baseline = {baseline_sim:>13,.0f} " + f"({fmt_pct((baseline_sim - observed) / observed)} vs observed)" + ) print(f" observed read calls = {observed_calls:>13,}") print(f" simulator calls (baseline) = {baseline_calls:>11,}") @@ -628,51 +775,77 @@ def report(result: dict, baseline_sim: float, baseline_calls: int, header = " D \\ L " + " ".join(f"{L:>6}" for L in line_caps) print(header) for i, D in enumerate(defaults): - row = " ".join(f"{grid[i,j]/baseline_sim:>6.2f}" for j in range(len(line_caps))) + row = " ".join( + f"{grid[i, j] / baseline_sim:>6.2f}" for j in range(len(line_caps)) + ) print(f" D={D:<6} {row}") - print(f"\nbest (D, L) = {result['best_DL']} → " - f"{grid[defaults.index(result['best_DL'][0]), line_caps.index(result['best_DL'][1])]:,.0f} tokens" - f" ({fmt_pct(grid.min()/baseline_sim - 1)})") + print( + f"\nbest (D, L) = {result['best_DL']} → " + f"{grid[defaults.index(result['best_DL'][0]), line_caps.index(result['best_DL'][1])]:,.0f} tokens" + f" ({fmt_pct(grid.min() / baseline_sim - 1)})" + ) # Summarizer threshold sweep at best (D, L). print(f"\nsummarizer threshold sweep at best (D, L) = {result['best_DL']}:") print(f" {'min_file_lines':<16} {'tokens':>12} {'vs baseline':>12}") for sm, t, k in result["summary_sweep"]: label = "off" if sm == -1 else ("always" if sm == 0 else f">={sm}") - print(f" {label:<16} {t:>12,.0f} {fmt_pct(t/baseline_sim - 1):>12}") - print(f"\nbest summarize_min = {result['best_summary'][0]} → " - f"{result['best_summary'][1]:,.0f} tokens " - f"({fmt_pct(result['best_summary'][1]/baseline_sim - 1)})") + print(f" {label:<16} {t:>12,.0f} {fmt_pct(t / baseline_sim - 1):>12}") + print( + f"\nbest summarize_min = {result['best_summary'][0]} → " + f"{result['best_summary'][1]:,.0f} tokens " + f"({fmt_pct(result['best_summary'][1] / baseline_sim - 1)})" + ) # Byte cap sweep at (best D, L, summarize_min). print(f"\nbyte cap sweep at best (D, L, summarize_min):") print(f" {'byte_cap':<10} {'tokens':>12} {'vs baseline':>12}") for bc, t, k in result["byte_cap_sweep"]: - print(f" {bc//1024:>4} KB {t:>12,.0f} {fmt_pct(t/baseline_sim - 1):>12}") - print(f"\nbest byte_cap = {result['best_byte_cap'][0]//1024} KB → " - f"{result['best_byte_cap'][1]:,.0f} tokens " - f"({fmt_pct(result['best_byte_cap'][1]/baseline_sim - 1)})") + print( + f" {bc // 1024:>4} KB {t:>12,.0f} {fmt_pct(t / baseline_sim - 1):>12}" + ) + print( + f"\nbest byte_cap = {result['best_byte_cap'][0] // 1024} KB → " + f"{result['best_byte_cap'][1]:,.0f} tokens " + f"({fmt_pct(result['best_byte_cap'][1] / baseline_sim - 1)})" + ) # Final recommendation. cfg = result["final_cfg"] print("\n" + "=" * 64) print(" RECOMMENDED CONFIG") print("=" * 64) - print(f" read.defaultLimit {cfg.default_page} lines (current: {CURRENT_DEFAULT})") + print( + f" read.defaultLimit {cfg.default_page} lines (current: {CURRENT_DEFAULT})" + ) print(f" read.lineCap {cfg.line_cap} lines (current: {CURRENT_LINE_CAP})") - print(f" read.byteCap {cfg.byte_cap//1024} KB (current: {CURRENT_BYTE_CAP//1024} KB)") - sm_label = "off" if cfg.summarize_min == -1 else ( - "always" if cfg.summarize_min == 0 else f"only files ≥ {cfg.summarize_min} lines") + print( + f" read.byteCap {cfg.byte_cap // 1024} KB (current: {CURRENT_BYTE_CAP // 1024} KB)" + ) + sm_label = ( + "off" + if cfg.summarize_min == -1 + else ( + "always" + if cfg.summarize_min == 0 + else f"only files ≥ {cfg.summarize_min} lines" + ) + ) print(f" read.summarizer {sm_label}") - print(f" simulated savings {fmt_pct(result['final_tokens']/baseline_sim - 1)} " - f"({baseline_sim - result['final_tokens']:,.0f} fewer tokens / window)") - print(f" calls {result['final_calls']:,} " - f"(baseline sim: {baseline_calls:,})") + print( + f" simulated savings {fmt_pct(result['final_tokens'] / baseline_sim - 1)} " + f"({baseline_sim - result['final_tokens']:,.0f} fewer tokens / window)" + ) + print( + f" calls {result['final_calls']:,} " + f"(baseline sim: {baseline_calls:,})" + ) # --------------------------------------------------------------------------- # # Entry + def main() -> int: ap = argparse.ArgumentParser(description=__doc__.splitlines()[1]) ap.add_argument("--since", default=DEFAULT_SINCE) @@ -694,10 +867,14 @@ def main() -> int: sizes = np.array([f.size_lines for f in files.values()], dtype=np.int64) tpls = np.array([f.tokens_per_line for f in files.values()], dtype=float) print(f" {len(files):,} distinct files") - print(f" file size p50={int(np.percentile(sizes,50))} " - f"p90={int(np.percentile(sizes,90))} max={int(sizes.max())}") - print(f" tokens/line p50={np.percentile(tpls,50):.2f} " - f"p90={np.percentile(tpls,90):.2f} max={tpls.max():.2f}") + print( + f" file size p50={int(np.percentile(sizes, 50))} " + f"p90={int(np.percentile(sizes, 90))} max={int(sizes.max())}" + ) + print( + f" tokens/line p50={np.percentile(tpls, 50):.2f} " + f"p90={np.percentile(tpls, 90):.2f} max={tpls.max():.2f}" + ) # Per-pair. by_pair = group_pairs(reads) @@ -705,8 +882,12 @@ def main() -> int: # Baseline simulation. print("\nsimulating baseline...") - baseline_cfg = Config(default_page=CURRENT_DEFAULT, line_cap=CURRENT_LINE_CAP, - byte_cap=CURRENT_BYTE_CAP, summarize_min=0) + baseline_cfg = Config( + default_page=CURRENT_DEFAULT, + line_cap=CURRENT_LINE_CAP, + byte_cap=CURRENT_BYTE_CAP, + summarize_min=0, + ) baseline_sim, baseline_calls = simulate(by_pair, files, baseline_cfg) observed, observed_calls = baseline_observed(reads) diff --git a/scripts/session-stats/plot_read_summarizer.py b/scripts/session-stats/plot_read_summarizer.py index 7b5325395..4b98130b2 100644 --- a/scripts/session-stats/plot_read_summarizer.py +++ b/scripts/session-stats/plot_read_summarizer.py @@ -15,6 +15,7 @@ thinking + user messages. That removes the "I worked harder that day" effect. Outputs to scripts/session-stats/out/read-summarizer-*.png. """ + from __future__ import annotations import argparse @@ -44,6 +45,7 @@ COHORT_COLORS = { # --------------------------------------------------------------------------- # # Classification + def has_selector(path: str) -> bool: """True iff `path` carries a read selector (`:50-200`, `:raw`, ...).""" if not path: @@ -76,6 +78,7 @@ def cohort_of(arg_json: str | None) -> str | None: # --------------------------------------------------------------------------- # # Data + def fetch_read_calls(conn) -> dict[str, dict[str, np.ndarray]]: sql = """ SELECT c.timestamp, @@ -97,7 +100,10 @@ def fetch_read_calls(conn) -> dict[str, dict[str, np.ndarray]]: out: dict[str, dict[str, np.ndarray]] = {} for c, rows in by.items(): if not rows: - out[c] = {"ts": np.array([], dtype=np.int64), "tok": np.array([], dtype=np.int64)} + out[c] = { + "ts": np.array([], dtype=np.int64), + "tok": np.array([], dtype=np.int64), + } continue ts = np.fromiter((r[0] for r in rows), dtype=np.int64, count=len(rows)) tok = np.fromiter((r[1] for r in rows), dtype=np.int64, count=len(rows)) @@ -152,7 +158,9 @@ def daily_sum(ts_ms: np.ndarray, tok: np.ndarray, day_axis: np.ndarray) -> np.nd return out -def daily_percentile(ts_ms: np.ndarray, tok: np.ndarray, q: float) -> tuple[np.ndarray, np.ndarray]: +def daily_percentile( + ts_ms: np.ndarray, tok: np.ndarray, q: float +) -> tuple[np.ndarray, np.ndarray]: if ts_ms.size == 0: return np.array([]), np.array([]) day_idx = ts_ms // DAY_MS @@ -164,7 +172,9 @@ def daily_percentile(ts_ms: np.ndarray, tok: np.ndarray, q: float) -> tuple[np.n lo, hi = order[i], order[i + 1] if hi > lo: pct[i] = np.percentile(tok[lo:hi], q) - dates = np.array([datetime.fromtimestamp(int(d) * DAY_MS / 1000, tz=timezone.utc) for d in days]) + dates = np.array( + [datetime.fromtimestamp(int(d) * DAY_MS / 1000, tz=timezone.utc) for d in days] + ) return dates, pct @@ -183,9 +193,10 @@ def smooth_nan(y: np.ndarray, w: int) -> np.ndarray: # --------------------------------------------------------------------------- # # Plot helpers + def thousands(x: float, _p=0) -> str: if x >= 1000: - return f"{x/1000:.1f}k" + return f"{x / 1000:.1f}k" return f"{x:.0f}" @@ -195,11 +206,20 @@ def style_time(ax: plt.Axes, deploy: datetime) -> None: ax.grid(True, alpha=0.25, linestyle="--") ax.axvline(deploy, color="#dc2626", linestyle="--", linewidth=1.2, alpha=0.8) y1 = ax.get_ylim()[1] if ax.get_ylim()[1] > 0 else 1 - ax.text(deploy, y1, " summarizer\n deploy", color="#dc2626", - va="top", ha="left", fontsize=9) + ax.text( + deploy, + y1, + " summarizer\n deploy", + color="#dc2626", + va="top", + ha="left", + fontsize=9, + ) -def panel_share_stacked(ax: plt.Axes, reads, denom_dates, denom, deploy: datetime) -> None: +def panel_share_stacked( + ax: plt.Axes, reads, denom_dates, denom, deploy: datetime +) -> None: """Stacked area: per-day read-cohort share of total tokens.""" series = [] labels = [] @@ -212,11 +232,13 @@ def panel_share_stacked(ax: plt.Axes, reads, denom_dates, denom, deploy: datetim series.append(smooth_nan(share, 7)) labels.append(cohort) colors.append(color) - x = np.array([datetime.fromtimestamp(int(d) / 1000, tz=timezone.utc) for d in denom_dates]) + x = np.array( + [datetime.fromtimestamp(int(d) / 1000, tz=timezone.utc) for d in denom_dates] + ) ax.stackplot(x, series, labels=labels, colors=colors, alpha=0.85) ax.set_title("read share of daily token spend (7d MA)") ax.set_ylabel("share of all tokens that day") - ax.yaxis.set_major_formatter(plt.FuncFormatter(lambda v, _: f"{v*100:.0f}%")) + ax.yaxis.set_major_formatter(plt.FuncFormatter(lambda v, _: f"{v * 100:.0f}%")) ax.set_ylim(0, None) ax.legend(loc="upper left", frameon=False) style_time(ax, deploy) @@ -224,7 +246,9 @@ def panel_share_stacked(ax: plt.Axes, reads, denom_dates, denom, deploy: datetim def panel_share_line(ax: plt.Axes, reads, denom_dates, denom, deploy: datetime) -> None: """Lines: each cohort's share, plus the combined total.""" - x = np.array([datetime.fromtimestamp(int(d) / 1000, tz=timezone.utc) for d in denom_dates]) + x = np.array( + [datetime.fromtimestamp(int(d) / 1000, tz=timezone.utc) for d in denom_dates] + ) total = np.zeros(denom_dates.size, dtype=np.int64) for cohort, color in COHORT_COLORS.items(): d = reads[cohort] @@ -235,22 +259,37 @@ def panel_share_line(ax: plt.Axes, reads, denom_dates, denom, deploy: datetime) ax.plot(x, smooth_nan(share, 7), label=cohort, color=color, linewidth=1.7) with np.errstate(divide="ignore", invalid="ignore"): combined = np.where(denom > 0, total / denom, 0.0) - ax.plot(x, smooth_nan(combined, 7), label="all reads", color="#111111", linewidth=2.2, linestyle="-") + ax.plot( + x, + smooth_nan(combined, 7), + label="all reads", + color="#111111", + linewidth=2.2, + linestyle="-", + ) ax.set_title("read share by cohort (7d MA)") ax.set_ylabel("share of daily tokens") - ax.yaxis.set_major_formatter(plt.FuncFormatter(lambda v, _: f"{v*100:.0f}%")) + ax.yaxis.set_major_formatter(plt.FuncFormatter(lambda v, _: f"{v * 100:.0f}%")) ax.set_ylim(0, None) ax.legend(loc="upper left", frameon=False) style_time(ax, deploy) -def panel_per_call(ax: plt.Axes, reads, deploy: datetime, q: float, label_q: str) -> None: +def panel_per_call( + ax: plt.Axes, reads, deploy: datetime, q: float, label_q: str +) -> None: for cohort, color in COHORT_COLORS.items(): d = reads[cohort] if d["ts"].size == 0: continue dates, pct = daily_percentile(d["ts"], d["tok"], q) - ax.plot(dates, smooth_nan(pct, 7), label=f"{cohort} ({label_q})", color=color, linewidth=1.9) + ax.plot( + dates, + smooth_nan(pct, 7), + label=f"{cohort} ({label_q})", + color=color, + linewidth=1.9, + ) ax.set_title(f"daily {label_q} tokens per read call (7d MA)") ax.set_ylabel("tokens / call") ax.set_yscale("log") @@ -262,14 +301,19 @@ def panel_per_call(ax: plt.Axes, reads, deploy: datetime, q: float, label_q: str # --------------------------------------------------------------------------- # # Stats + def share_stats(reads, denom_dates, denom, deploy_ms: int) -> None: pre_mask = denom_dates < deploy_ms post_mask = denom_dates >= deploy_ms pre_total = int(denom[pre_mask].sum()) post_total = int(denom[post_mask].sum()) print(f"\nshare-of-day (pre vs post deploy):") - print(f" denominator pre = {pre_total:>14,} tokens across {int(pre_mask.sum())} days") - print(f" denominator post = {post_total:>14,} tokens across {int(post_mask.sum())} days") + print( + f" denominator pre = {pre_total:>14,} tokens across {int(pre_mask.sum())} days" + ) + print( + f" denominator post = {post_total:>14,} tokens across {int(post_mask.sum())} days" + ) print(f" {'cohort':<22} {'pre share':>10} {'post share':>11} {'delta':>10}") grand_pre = 0 grand_post = 0 @@ -282,17 +326,23 @@ def share_stats(reads, denom_dates, denom, deploy_ms: int) -> None: grand_post += post pre_share = pre / pre_total if pre_total else 0 post_share = post / post_total if post_total else 0 - print(f" {cohort:<22} {pre_share*100:>9.2f}% {post_share*100:>10.2f}% " - f"{(post_share-pre_share)*100:>+9.2f}pp") + print( + f" {cohort:<22} {pre_share * 100:>9.2f}% {post_share * 100:>10.2f}% " + f"{(post_share - pre_share) * 100:>+9.2f}pp" + ) pre_share = grand_pre / pre_total if pre_total else 0 post_share = grand_post / post_total if post_total else 0 - print(f" {'all reads':<22} {pre_share*100:>9.2f}% {post_share*100:>10.2f}% " - f"{(post_share-pre_share)*100:>+9.2f}pp") + print( + f" {'all reads':<22} {pre_share * 100:>9.2f}% {post_share * 100:>10.2f}% " + f"{(post_share - pre_share) * 100:>+9.2f}pp" + ) def per_call_stats(reads, deploy_ms: int) -> None: print(f"\nper-call stats (pre vs post deploy):") - print(f" {'cohort':<22} {'window':<6} {'n':>9} {'p50':>7} {'p90':>7} {'mean':>8}") + print( + f" {'cohort':<22} {'window':<6} {'n':>9} {'p50':>7} {'p90':>7} {'mean':>8}" + ) for cohort in COHORT_COLORS: d = reads[cohort] if d["ts"].size == 0: @@ -302,19 +352,25 @@ def per_call_stats(reads, deploy_ms: int) -> None: for name, arr in (("pre", pre), ("post", post)): if arr.size == 0: continue - print(f" {cohort:<22} {name:<6} {arr.size:>9,} " - f"{int(np.percentile(arr,50)):>7,} " - f"{int(np.percentile(arr,90)):>7,} " - f"{int(arr.mean()):>8,}") + print( + f" {cohort:<22} {name:<6} {arr.size:>9,} " + f"{int(np.percentile(arr, 50)):>7,} " + f"{int(np.percentile(arr, 90)):>7,} " + f"{int(arr.mean()):>8,}" + ) # --------------------------------------------------------------------------- # # Entry + def main() -> int: ap = argparse.ArgumentParser(description="read summarizer impact analysis") - ap.add_argument("--deploy", default=DEFAULT_DEPLOY, - help=f"deploy date YYYY-MM-DD (default {DEFAULT_DEPLOY})") + ap.add_argument( + "--deploy", + default=DEFAULT_DEPLOY, + help=f"deploy date YYYY-MM-DD (default {DEFAULT_DEPLOY})", + ) args = ap.parse_args() deploy = datetime.strptime(args.deploy, "%Y-%m-%d").replace(tzinfo=timezone.utc) @@ -343,7 +399,9 @@ def main() -> int: panel_share_line(axes[0, 1], reads, denom_dates, denom, deploy) panel_per_call(axes[1, 0], reads, deploy, q=50, label_q="p50") panel_per_call(axes[1, 1], reads, deploy, q=90, label_q="p90") - fig.suptitle(f"read summarizer impact — deploy = {args.deploy}", fontsize=13, y=0.995) + fig.suptitle( + f"read summarizer impact — deploy = {args.deploy}", fontsize=13, y=0.995 + ) fig.tight_layout() p = OUT_DIR / "read-summarizer.png" fig.savefig(p, bbox_inches="tight") diff --git a/scripts/session-stats/plot_tools.py b/scripts/session-stats/plot_tools.py index 82453dd91..363be51ac 100644 --- a/scripts/session-stats/plot_tools.py +++ b/scripts/session-stats/plot_tools.py @@ -15,6 +15,7 @@ N by total tokens (default 10); override with --top N or --tools a,b,c. Output: scripts/session-stats/out/tool-trends.png + standalone panels. """ + from __future__ import annotations import argparse @@ -39,9 +40,18 @@ TOOL_ALIAS = {"grep": "search"} # 10-class qualitative palette (tab10) — distinct hues for line + area work. PALETTE = [ - "#1f77b4", "#d62728", "#2ca02c", "#ff7f0e", "#9467bd", - "#8c564b", "#17becf", "#e377c2", "#bcbd22", "#7f7f7f", - "#393b79", "#637939", + "#1f77b4", + "#d62728", + "#2ca02c", + "#ff7f0e", + "#9467bd", + "#8c564b", + "#17becf", + "#e377c2", + "#bcbd22", + "#7f7f7f", + "#393b79", + "#637939", ] @@ -56,6 +66,7 @@ def normalize_case_sql(col: str) -> str: # --------------------------------------------------------------------------- # # Data access + def _connect() -> sqlite3.Connection: if not DB_PATH.exists(): sys.exit(f"db missing: {DB_PATH}") @@ -148,7 +159,10 @@ def fetch_per_call(conn: sqlite3.Connection, tools: list[str]) -> dict[str, dict out: dict[str, dict] = {} for t, rows in by_tool.items(): if not rows: - out[t] = {"ts": np.array([], dtype=np.int64), "tok": np.array([], dtype=np.int64)} + out[t] = { + "ts": np.array([], dtype=np.int64), + "tok": np.array([], dtype=np.int64), + } continue ts = np.fromiter((r[0] for r in rows), dtype=np.int64, count=len(rows)) tok = np.fromiter((r[1] for r in rows), dtype=np.int64, count=len(rows)) @@ -160,6 +174,7 @@ def fetch_per_call(conn: sqlite3.Connection, tools: list[str]) -> dict[str, dict # --------------------------------------------------------------------------- # # Helpers + def smooth(y: np.ndarray, w: int = 7) -> np.ndarray: if w <= 1 or len(y) < w: return y.astype(float) @@ -208,7 +223,10 @@ def weekly_median(ts_ms: np.ndarray, tok: np.ndarray) -> tuple[np.ndarray, np.nd if hi > lo: p50[i] = np.percentile(tok[lo:hi], 50) week_dates = np.array( - [datetime.fromtimestamp(int(w) * WEEK_MS / 1000, tz=timezone.utc) for w in weeks] + [ + datetime.fromtimestamp(int(w) * WEEK_MS / 1000, tz=timezone.utc) + for w in weeks + ] ) return week_dates, p50 @@ -216,10 +234,15 @@ def weekly_median(ts_ms: np.ndarray, tok: np.ndarray) -> tuple[np.ndarray, np.nd # --------------------------------------------------------------------------- # # Panels -def panel_total_tokens(ax: plt.Axes, daily: dict, tools: list[str], colors: dict) -> None: + +def panel_total_tokens( + ax: plt.Axes, daily: dict, tools: list[str], colors: dict +) -> None: dates = daily["dates"] series = [smooth(daily[t]["args"] + daily[t]["results"]) for t in tools] - ax.stackplot(dates, series, labels=tools, colors=[colors[t] for t in tools], alpha=0.9) + ax.stackplot( + dates, series, labels=tools, colors=[colors[t] for t in tools], alpha=0.9 + ) ax.set_title("Daily token volume (args + results, 7d MA)") ax.set_ylabel("tokens / day") ax.yaxis.set_major_formatter(plt.FuncFormatter(millions)) @@ -227,17 +250,23 @@ def panel_total_tokens(ax: plt.Axes, daily: dict, tools: list[str], colors: dict style_time_axis(ax) -def panel_call_counts(ax: plt.Axes, daily: dict, tools: list[str], colors: dict) -> None: +def panel_call_counts( + ax: plt.Axes, daily: dict, tools: list[str], colors: dict +) -> None: dates = daily["dates"] for t in tools: - ax.plot(dates, smooth(daily[t]["calls"]), label=t, color=colors[t], linewidth=1.6) + ax.plot( + dates, smooth(daily[t]["calls"]), label=t, color=colors[t], linewidth=1.6 + ) ax.set_title("Daily call count (7d MA)") ax.set_ylabel("calls / day") ax.legend(loc="upper left", frameon=False, ncol=2, fontsize=9) style_time_axis(ax) -def panel_mean_per_call(ax: plt.Axes, daily: dict, tools: list[str], colors: dict) -> None: +def panel_mean_per_call( + ax: plt.Axes, daily: dict, tools: list[str], colors: dict +) -> None: dates = daily["dates"] for t in tools: totals = daily[t]["args"] + daily[t]["results"] @@ -265,7 +294,9 @@ def panel_cumulative(ax: plt.Axes, daily: dict, tools: list[str], colors: dict) style_time_axis(ax) -def panel_weekly_median(ax: plt.Axes, per_call: dict, tools: list[str], colors: dict) -> None: +def panel_weekly_median( + ax: plt.Axes, per_call: dict, tools: list[str], colors: dict +) -> None: for t in tools: w, p50 = weekly_median(per_call[t]["ts"], per_call[t]["tok"]) if w.size == 0: @@ -279,8 +310,12 @@ def panel_weekly_median(ax: plt.Axes, per_call: dict, tools: list[str], colors: style_time_axis(ax) -def panel_histogram(ax: plt.Axes, per_call: dict, tools: list[str], colors: dict) -> None: - all_tok = np.concatenate([per_call[t]["tok"] for t in tools if per_call[t]["tok"].size]) +def panel_histogram( + ax: plt.Axes, per_call: dict, tools: list[str], colors: dict +) -> None: + all_tok = np.concatenate( + [per_call[t]["tok"] for t in tools if per_call[t]["tok"].size] + ) if all_tok.size == 0: return hi = max(all_tok.max(), 10) @@ -312,6 +347,7 @@ def panel_histogram(ax: plt.Axes, per_call: dict, tools: list[str], colors: dict # --------------------------------------------------------------------------- # # Entry + def main() -> int: ap = argparse.ArgumentParser(description=__doc__.splitlines()[1]) ap.add_argument("--top", type=int, default=10, help="top N tools by total tokens") @@ -364,12 +400,12 @@ def main() -> int: print(f"wrote {combined}") panels: tuple[tuple[str, Callable, dict], ...] = ( - ("daily-tokens", panel_total_tokens, daily), - ("daily-calls", panel_call_counts, daily), - ("tokens-per-call", panel_mean_per_call, daily), - ("cumulative-tokens", panel_cumulative, daily), - ("per-call-median", panel_weekly_median, per_call), - ("per-call-histogram", panel_histogram, per_call), + ("daily-tokens", panel_total_tokens, daily), + ("daily-calls", panel_call_counts, daily), + ("tokens-per-call", panel_mean_per_call, daily), + ("cumulative-tokens", panel_cumulative, daily), + ("per-call-median", panel_weekly_median, per_call), + ("per-call-histogram", panel_histogram, per_call), ) for name, fn, src in panels: f2, ax = plt.subplots(figsize=(11, 5)) diff --git a/scripts/session-stats/read_optimizer.py b/scripts/session-stats/read_optimizer.py index d18d45da4..02888025f 100644 --- a/scripts/session-stats/read_optimizer.py +++ b/scripts/session-stats/read_optimizer.py @@ -20,6 +20,7 @@ truncations, and a Pareto frontier. Output: scripts/session-stats/out/read-optimizer.png """ + from __future__ import annotations import argparse @@ -51,12 +52,54 @@ READ_MAX_COLUMN = 768 _RANGE_RE = re.compile(r"^(\d+)(?:([-+])(\d+))?$") TEXT_EXTS = { - ".ts", ".tsx", ".js", ".jsx", ".mts", ".cts", ".mjs", ".cjs", - ".rs", ".go", ".py", ".rb", ".java", ".kt", ".kts", ".c", ".cc", - ".cpp", ".h", ".hpp", ".cs", ".swift", ".php", ".lua", ".sh", - ".bash", ".zsh", ".fish", ".md", ".txt", ".json", ".jsonc", ".json5", - ".yaml", ".yml", ".toml", ".xml", ".html", ".css", ".scss", ".sql", - ".adoc", ".typ", ".rsx", ".vue", ".svelte", ".dockerfile", "", + ".ts", + ".tsx", + ".js", + ".jsx", + ".mts", + ".cts", + ".mjs", + ".cjs", + ".rs", + ".go", + ".py", + ".rb", + ".java", + ".kt", + ".kts", + ".c", + ".cc", + ".cpp", + ".h", + ".hpp", + ".cs", + ".swift", + ".php", + ".lua", + ".sh", + ".bash", + ".zsh", + ".fish", + ".md", + ".txt", + ".json", + ".jsonc", + ".json5", + ".yaml", + ".yml", + ".toml", + ".xml", + ".html", + ".css", + ".scss", + ".sql", + ".adoc", + ".typ", + ".rsx", + ".vue", + ".svelte", + ".dockerfile", + "", } @@ -65,7 +108,7 @@ class ReadCall: session: str file: str seq: int - kind: str # explicit | open | default | raw | conflicts | other + kind: str # explicit | open | default | raw | conflicts | other start: int | None end: int | None arg_tokens: int @@ -179,7 +222,12 @@ def parse_call(row) -> ReadCall | None: if kind == "default": offset = obj.get("offset") limit = obj.get("limit") - if isinstance(offset, int) and offset >= 1 and isinstance(limit, int) and limit >= 1: + if ( + isinstance(offset, int) + and offset >= 1 + and isinstance(limit, int) + and limit >= 1 + ): kind = "explicit" start = offset end = offset + limit - 1 @@ -258,7 +306,9 @@ def is_covered(intervals: list[tuple[int, int]], target: tuple[int, int]) -> boo return False -def add_interval(intervals: list[tuple[int, int]], item: tuple[int, int]) -> list[tuple[int, int]]: +def add_interval( + intervals: list[tuple[int, int]], item: tuple[int, int] +) -> list[tuple[int, int]]: s, e = item out: list[tuple[int, int]] = [] placed = False @@ -278,7 +328,9 @@ def add_interval(intervals: list[tuple[int, int]], item: tuple[int, int]) -> lis return out -def estimate_cost(call: ReadCall, delivered: tuple[int, int]) -> tuple[float, bool, bool]: +def estimate_cost( + call: ReadCall, delivered: tuple[int, int] +) -> tuple[float, bool, bool]: lines = max(0, delivered[1] - delivered[0] + 1) line_tokens = call.token_per_line * lines # Approximate byte cap. The implementation scales byte cap as @@ -289,10 +341,16 @@ def estimate_cost(call: ReadCall, delivered: tuple[int, int]) -> tuple[float, bo bytes_limited = approx_bytes > byte_budget if bytes_limited: line_tokens = byte_budget / 4 - return call.arg_tokens + line_tokens, call.kind == "explicit" and lines >= call.config_max_lines if False else False, bytes_limited + return ( + call.arg_tokens + line_tokens, + call.kind == "explicit" and lines >= call.config_max_lines if False else False, + bytes_limited, + ) -def load_reads(conn: sqlite3.Connection, since_ms: int) -> dict[tuple[str, str], list[ReadCall]]: +def load_reads( + conn: sqlite3.Connection, since_ms: int +) -> dict[tuple[str, str], list[ReadCall]]: sql = """ SELECT c.session_file, c.seq, c.arg_json, COALESCE(c.arg_tokens,0), COALESCE(r.result_tokens,0) @@ -372,7 +430,9 @@ def replay(groups: dict[tuple[str, str], list[ReadCall]], cfg: Config) -> Replay lines = delivered[1] - delivered[0] + 1 if lines >= cfg.max_lines and call.kind == "explicit": # Candidate max cap would truncate this explicit request. - requested_len = max(1, (call.end or call.start or 1) - (call.start or 1) + 1) + requested_len = max( + 1, (call.end or call.start or 1) - (call.start or 1) + 1 + ) if requested_len + cfg.leading + cfg.trailing > cfg.max_lines: trunc += 1 line_tokens = call.token_per_line * lines @@ -443,10 +503,18 @@ def candidate_grid(args) -> list[Config]: return out -def pareto(results: list[ReplayResult], max_truncations: int, max_regret_tokens: float = math.inf) -> list[ReplayResult]: +def pareto( + results: list[ReplayResult], + max_truncations: int, + max_regret_tokens: float = math.inf, +) -> list[ReplayResult]: # Frontier over (tokens lower, calls lower), excluding configs that truncate # more explicit requests than today's cap. - clean = [r for r in results if r.truncations <= max_truncations and r.tokens <= max_regret_tokens] + clean = [ + r + for r in results + if r.truncations <= max_truncations and r.tokens <= max_regret_tokens + ] clean.sort(key=lambda r: (r.tokens, r.calls)) frontier: list[ReplayResult] = [] best_calls = math.inf @@ -457,13 +525,16 @@ def pareto(results: list[ReplayResult], max_truncations: int, max_regret_tokens: return frontier -def choose_recommended(results: list[ReplayResult], current: ReplayResult) -> ReplayResult: +def choose_recommended( + results: list[ReplayResult], current: ReplayResult +) -> ReplayResult: # Objective: minimize tokens plus a small penalty for still needing calls, # while requiring no *additional* explicit-request truncations and at least # current first-call coverage. One avoided read call is valued at ~250 # tokens of ergonomics. viable = [ - r for r in results + r + for r in results if r.truncations <= current.truncations and r.first_cover_rate >= current.first_cover_rate and r.tokens <= current.tokens * 1.02 @@ -472,7 +543,14 @@ def choose_recommended(results: list[ReplayResult], current: ReplayResult) -> Re viable = [r for r in results if r.truncations <= current.truncations] if not viable: viable = results - return min(viable, key=lambda r: r.tokens + 250 * r.calls + 100_000 * max(0, r.truncations - current.truncations)) + return min( + viable, + key=lambda r: ( + r.tokens + + 250 * r.calls + + 100_000 * max(0, r.truncations - current.truncations) + ), + ) def print_result(prefix: str, r: ReplayResult, baseline: ReplayResult) -> None: @@ -480,15 +558,17 @@ def print_result(prefix: str, r: ReplayResult, baseline: ReplayResult) -> None: dcalls = r.calls - baseline.calls print( f"{prefix:<14} {r.config.label():<22} " - f"tokens={r.tokens/1e6:8.2f}M ({dtok/baseline.tokens*100:+6.2f}%) " + f"tokens={r.tokens / 1e6:8.2f}M ({dtok / baseline.tokens * 100:+6.2f}%) " f"calls={r.calls:7,} ({dcalls:+7,}) " f"skipped={r.skipped_calls:6,} " - f"first-cover={r.first_cover_rate*100:5.1f}% " + f"first-cover={r.first_cover_rate * 100:5.1f}% " f"trunc={r.truncations:4,}" ) -def plot(results: list[ReplayResult], current: ReplayResult, recommended: ReplayResult) -> Path: +def plot( + results: list[ReplayResult], current: ReplayResult, recommended: ReplayResult +) -> Path: OUT_DIR.mkdir(parents=True, exist_ok=True) plt.rcParams.update({"figure.dpi": 110, "font.size": 10}) fig, axes = plt.subplots(2, 2, figsize=(15, 9)) @@ -499,9 +579,25 @@ def plot(results: list[ReplayResult], current: ReplayResult, recommended: Replay sizes = np.array([20 + min(80, r.config.trailing * 5) for r in results]) ax = axes[0, 0] - sc = ax.scatter(xs, ys, c=colors, s=sizes, cmap="viridis", alpha=0.65, edgecolors="none") - ax.scatter([current.calls], [current.tokens / 1e6], marker="*", s=180, color="#111", label="current") - ax.scatter([recommended.calls], [recommended.tokens / 1e6], marker="*", s=180, color="#dc2626", label="recommended") + sc = ax.scatter( + xs, ys, c=colors, s=sizes, cmap="viridis", alpha=0.65, edgecolors="none" + ) + ax.scatter( + [current.calls], + [current.tokens / 1e6], + marker="*", + s=180, + color="#111", + label="current", + ) + ax.scatter( + [recommended.calls], + [recommended.tokens / 1e6], + marker="*", + s=180, + color="#dc2626", + label="recommended", + ) ax.set_xlabel("paid read calls after replay") ax.set_ylabel("estimated read tokens (M)") ax.set_title("candidate trade-off: tokens vs follow-up calls") @@ -513,10 +609,34 @@ def plot(results: list[ReplayResult], current: ReplayResult, recommended: Replay ax = axes[0, 1] frontier = pareto(results, current.truncations) frontier.sort(key=lambda r: r.calls) - ax.plot([r.calls for r in frontier], [r.tokens / 1e6 for r in frontier], color="#2563eb", linewidth=2) - ax.scatter([r.calls for r in frontier], [r.tokens / 1e6 for r in frontier], color="#2563eb", s=20) - ax.scatter([current.calls], [current.tokens / 1e6], marker="*", s=180, color="#111", label="current") - ax.scatter([recommended.calls], [recommended.tokens / 1e6], marker="*", s=180, color="#dc2626", label="recommended") + ax.plot( + [r.calls for r in frontier], + [r.tokens / 1e6 for r in frontier], + color="#2563eb", + linewidth=2, + ) + ax.scatter( + [r.calls for r in frontier], + [r.tokens / 1e6 for r in frontier], + color="#2563eb", + s=20, + ) + ax.scatter( + [current.calls], + [current.tokens / 1e6], + marker="*", + s=180, + color="#111", + label="current", + ) + ax.scatter( + [recommended.calls], + [recommended.tokens / 1e6], + marker="*", + s=180, + color="#dc2626", + label="recommended", + ) ax.set_xlabel("paid read calls") ax.set_ylabel("estimated read tokens (M)") ax.set_title("Pareto frontier (no extra explicit truncations)") @@ -526,20 +646,31 @@ def plot(results: list[ReplayResult], current: ReplayResult, recommended: Replay ax = axes[1, 0] by_default: dict[int, list[ReplayResult]] = defaultdict(list) for r in results: - if r.truncations <= current.truncations and r.config.leading == recommended.config.leading and r.config.trailing == recommended.config.trailing: + if ( + r.truncations <= current.truncations + and r.config.leading == recommended.config.leading + and r.config.trailing == recommended.config.trailing + ): by_default[r.config.default].append(r) defaults = sorted(by_default) vals = [min(v, key=lambda r: r.tokens).tokens / 1e6 for v in by_default.values()] ax.bar([str(d) for d in defaults], vals, color="#16a34a") - ax.axhline(current.tokens / 1e6, color="#111", linestyle="--", linewidth=1, label="current") + ax.axhline( + current.tokens / 1e6, color="#111", linestyle="--", linewidth=1, label="current" + ) ax.set_xlabel("defaultLimit") ax.set_ylabel("best tokens (M)") - ax.set_title(f"defaultLimit sensitivity (L={recommended.config.leading}, T={recommended.config.trailing})") + ax.set_title( + f"defaultLimit sensitivity (L={recommended.config.leading}, T={recommended.config.trailing})" + ) ax.legend(frameon=False) ax.grid(True, axis="y", alpha=0.25, linestyle="--") ax = axes[1, 1] - top = sorted([r for r in results if r.truncations <= current.truncations], key=lambda r: r.tokens + 250 * r.calls)[:12] + top = sorted( + [r for r in results if r.truncations <= current.truncations], + key=lambda r: r.tokens + 250 * r.calls, + )[:12] labels = [r.config.label() for r in top] token_delta = [(r.tokens - current.tokens) / current.tokens * 100 for r in top] call_delta = [(r.calls - current.calls) / current.calls * 100 for r in top] @@ -564,7 +695,9 @@ def plot(results: list[ReplayResult], current: ReplayResult, recommended: Replay def main() -> int: ap = argparse.ArgumentParser(description="read configuration optimizer") - ap.add_argument("--since", default=DEFAULT_SINCE, help=f"YYYY-MM-DD (default {DEFAULT_SINCE})") + ap.add_argument( + "--since", default=DEFAULT_SINCE, help=f"YYYY-MM-DD (default {DEFAULT_SINCE})" + ) ap.add_argument("--defaults", default="100,150,200,250,300,400,500,700,1000") ap.add_argument("--max-lines", default="500,750,1000,1500,2000,3000") ap.add_argument("--leading", default="0,3,5,10,20") @@ -581,12 +714,19 @@ def main() -> int: groups = load_reads(conn, since_ms) conn.close() total_calls = sum(len(v) for v in groups.values()) - print(f"loaded {total_calls:,} read calls across {len(groups):,} (session,file) groups since {args.since}") + print( + f"loaded {total_calls:,} read calls across {len(groups):,} (session,file) groups since {args.since}" + ) - current = replay(groups, Config(CURRENT_DEFAULT, CURRENT_MAX_LINES, CURRENT_LEADING, CURRENT_TRAILING)) + current = replay( + groups, + Config(CURRENT_DEFAULT, CURRENT_MAX_LINES, CURRENT_LEADING, CURRENT_TRAILING), + ) configs = candidate_grid(args) # Ensure current is present even if user overrides grid. - cur_cfg = Config(CURRENT_DEFAULT, CURRENT_MAX_LINES, CURRENT_LEADING, CURRENT_TRAILING) + cur_cfg = Config( + CURRENT_DEFAULT, CURRENT_MAX_LINES, CURRENT_LEADING, CURRENT_TRAILING + ) if cur_cfg not in configs: configs.append(cur_cfg) print(f"evaluating {len(configs):,} candidate configs") @@ -598,22 +738,32 @@ def main() -> int: print_result("recommended", recommended, current) allowed = [r for r in results if r.truncations <= current.truncations] - print(f"\nTop token-minimizing configs (truncations <= current {current.truncations:,}):") + print( + f"\nTop token-minimizing configs (truncations <= current {current.truncations:,}):" + ) for i, r in enumerate(sorted(allowed, key=lambda r: r.tokens)[: args.top], 1): print_result(f"#{i}", r, current) - print(f"\nTop balanced configs (tokens + 250 tokens/read-call objective, truncations <= current {current.truncations:,}):") - for i, r in enumerate(sorted(allowed, key=lambda r: r.tokens + 250 * r.calls)[: args.top], 1): + print( + f"\nTop balanced configs (tokens + 250 tokens/read-call objective, truncations <= current {current.truncations:,}):" + ) + for i, r in enumerate( + sorted(allowed, key=lambda r: r.tokens + 250 * r.calls)[: args.top], 1 + ): print_result(f"#{i}", r, current) no_call_increase = [r for r in allowed if r.calls <= current.calls] - print(f"\nBest configs with calls <= current (truncations <= current {current.truncations:,}):") - for i, r in enumerate(sorted(no_call_increase, key=lambda r: r.tokens)[: args.top], 1): + print( + f"\nBest configs with calls <= current (truncations <= current {current.truncations:,}):" + ) + for i, r in enumerate( + sorted(no_call_increase, key=lambda r: r.tokens)[: args.top], 1 + ): print_result(f"#{i}", r, current) print("\nRecommended breakdown:") print(f" selector groups : {recommended.selector_groups:,}") - print(f" selector first-cover : {recommended.first_cover_rate*100:.1f}%") + print(f" selector first-cover : {recommended.first_cover_rate * 100:.1f}%") print(f" selector skipped calls : {recommended.selector_skipped:,}") print(f" default skipped calls : {recommended.default_skipped:,}") print(f" raw/unmodelled calls : {recommended.raw_calls:,}") diff --git a/scripts/session-stats/sync.py b/scripts/session-stats/sync.py index d50132939..31229ef38 100644 --- a/scripts/session-stats/sync.py +++ b/scripts/session-stats/sync.py @@ -251,8 +251,12 @@ def batch_count_tokens(strings: list[str]) -> list[int]: _HEADER_NEW_RE = re.compile(r"^¶+\s*([^\s#¶]+)(?:#\S+)?\s*$") # Verb-based v4 (current) ops; body rows are `+TEXT` on the following lines. -_VERB_REPLACE_RE = re.compile(r"^\s*replace\s+([1-9][0-9]*)(?:\s*(?:\.\.|-|…)\s*([1-9][0-9]*))?\s*:?\s*$") -_VERB_DELETE_RE = re.compile(r"^\s*delete\s+([1-9][0-9]*)(?:\s*(?:\.\.|-|…)\s*([1-9][0-9]*))?\s*$") +_VERB_REPLACE_RE = re.compile( + r"^\s*replace\s+([1-9][0-9]*)(?:\s*(?:\.\.|-|…)\s*([1-9][0-9]*))?\s*:?\s*$" +) +_VERB_DELETE_RE = re.compile( + r"^\s*delete\s+([1-9][0-9]*)(?:\s*(?:\.\.|-|…)\s*([1-9][0-9]*))?\s*$" +) _VERB_INSERT_RE = re.compile( r"^\s*insert\s+(?:(?Pbefore|after)\s+(?P[1-9][0-9]*)|(?Phead|tail))\s*:?\s*$" ) @@ -324,9 +328,9 @@ class EditSection: def parse_hashline_input(input_str: str) -> list[EditSection]: sections: list[EditSection] = [] cur: EditSection | None = None - cur_format: str | None = None # "hash" (¶) | "legacy" (§) + cur_format: str | None = None # "hash" (¶) | "legacy" (§) cur_grammar: str | None = None # within "hash": None | "verb" | "sigil" - open_idx: int | None = None # current open payload block in cur + open_idx: int | None = None # current open payload block in cur def open_new(s: EditSection) -> int: s.payload_blocks.append([]) @@ -581,6 +585,7 @@ def extract_warnings(text: str) -> list[str]: # --------------------------------------------------------------------------- # # JSONL parsing + def parse_iso_ms(s: str | None) -> int: if not s: return 0 @@ -588,6 +593,7 @@ def parse_iso_ms(s: str | None) -> int: if s.endswith("Z"): s = s[:-1] + "+00:00" from datetime import datetime + return int(datetime.fromisoformat(s).timestamp() * 1000) except Exception: return 0 @@ -633,10 +639,12 @@ class SessionRecords: tool_results: list[list] = field(default_factory=list) assistant_msgs: list[list] = field(default_factory=list) user_msgs: list[list] = field(default_factory=list) - edit_calls: list[tuple] = field(default_factory=list) # initial stub on toolCall - edit_call_results: list[tuple] = field(default_factory=list) # success+warnings on toolResult - edit_sections: list[tuple] = field(default_factory=list) # one row per section - pending_tokens: list[tuple] = field(default_factory=list) # (row, field_idx, text) + edit_calls: list[tuple] = field(default_factory=list) # initial stub on toolCall + edit_call_results: list[tuple] = field( + default_factory=list + ) # success+warnings on toolResult + edit_sections: list[tuple] = field(default_factory=list) # one row per section + pending_tokens: list[tuple] = field(default_factory=list) # (row, field_idx, text) starting_seq: int = 0 full_rebuild: bool = False starting_offset: int = 0 @@ -759,11 +767,21 @@ def _ingest_assistant(rec, path, seq, entry_id, ts, msg, content) -> None: elif isinstance(arg_obj, str): arg_json = arg_obj else: - arg_json = json.dumps(arg_obj, separators=(",", ":"), ensure_ascii=False) + arg_json = json.dumps( + arg_obj, separators=(",", ":"), ensure_ascii=False + ) row = [ - sf, seq, entry_id, call_id, - tool_name, raw_name, ts, model, provider, - arg_json, 0, + sf, + seq, + entry_id, + call_id, + tool_name, + raw_name, + ts, + model, + provider, + arg_json, + 0, ] rec.tool_calls.append(row) if arg_json: @@ -785,8 +803,16 @@ def _ingest_assistant(rec, path, seq, entry_id, ts, msg, content) -> None: thinking_tokens_slot = 0 if text_blob or thinking_blob: row = [ - sf, seq, entry_id, ts, model, provider, - text_blob, thinking_blob, text_tokens_slot, thinking_tokens_slot, + sf, + seq, + entry_id, + ts, + model, + provider, + text_blob, + thinking_blob, + text_tokens_slot, + thinking_tokens_slot, ] rec.assistant_msgs.append(row) if text_blob: @@ -816,11 +842,11 @@ def _ingest_edit_call(rec, sf, seq, ts, call_id, arg_obj, arg_json) -> None: raw_input_len = len(input_str.encode("utf-8")) # Stub call row (success + warnings come from toolResult later). - rec.edit_calls.append( - (sf, call_id, seq, ts, raw_input_len, EDIT_PARSER_VERSION) - ) + rec.edit_calls.append((sf, call_id, seq, ts, raw_input_len, EDIT_PARSER_VERSION)) - if not any(line.startswith(("¶", "§")) for line in input_str.lstrip("\ufeff").splitlines()): + if not any( + line.startswith(("¶", "§")) for line in input_str.lstrip("\ufeff").splitlines() + ): # Vim-mode or other shape — no sections to record. return @@ -849,12 +875,22 @@ def _ingest_edit_call(rec, sf, seq, ts, call_id, arg_obj, arg_json) -> None: rec.edit_sections.append( ( - sf, call_id, seq, idx, sec.target_file, - sec.op_count, sec.deleted_lines, sec.payload_count, sec.change_size, - sec.min_line, sec.max_line, + sf, + call_id, + seq, + idx, + sec.target_file, + sec.op_count, + sec.deleted_lines, + sec.payload_count, + sec.change_size, + sec.min_line, + sec.max_line, json.dumps(sec.payload_blocks, ensure_ascii=False), json.dumps(sec.op_anchors, ensure_ascii=False), - longest_repeat_len, repeat_block_idx, sample, + longest_repeat_len, + repeat_block_idx, + sample, json.dumps(dups, ensure_ascii=False), EDIT_PARSER_VERSION, ) @@ -893,6 +929,7 @@ def _ingest_user(rec, path, seq, entry_id, ts, content) -> None: # --------------------------------------------------------------------------- # # DB + def open_db() -> sqlite3.Connection: DB_PATH.parent.mkdir(parents=True, exist_ok=True) conn = sqlite3.connect(DB_PATH, isolation_level=None, check_same_thread=False) @@ -905,7 +942,9 @@ def open_db() -> sqlite3.Connection: return conn -def existing_state(conn: sqlite3.Connection) -> dict[str, tuple[int, int, int, int, int]]: +def existing_state( + conn: sqlite3.Connection, +) -> dict[str, tuple[int, int, int, int, int]]: """{session_file: (mtime, size, byte_offset, line_count, parser_version)}""" rows = conn.execute( "SELECT session_file, mtime, size, byte_offset, line_count, parser_version " @@ -921,9 +960,12 @@ def write_records(conn: sqlite3.Connection, rec: SessionRecords, now_ms: int) -> try: if rec.full_rebuild: for tbl in ( - "ss_tool_calls", "ss_tool_results", - "ss_assistant_msgs", "ss_user_msgs", - "ss_edit_calls", "ss_edit_sections", + "ss_tool_calls", + "ss_tool_results", + "ss_assistant_msgs", + "ss_user_msgs", + "ss_edit_calls", + "ss_edit_sections", ): cur.execute(f"DELETE FROM {tbl} WHERE session_file = ?", (sf,)) @@ -979,8 +1021,10 @@ def write_records(conn: sqlite3.Connection, rec: SessionRecords, now_ms: int) -> "ON CONFLICT(session_file, call_id) DO UPDATE SET " " success=excluded.success, warnings=excluded.warnings, " " parser_version=excluded.parser_version", - [(sf_, cid, succ, warn, EDIT_PARSER_VERSION) - for (sf_, cid, succ, warn) in rec.edit_call_results], + [ + (sf_, cid, succ, warn, EDIT_PARSER_VERSION) + for (sf_, cid, succ, warn) in rec.edit_call_results + ], ) if rec.edit_sections: cur.executemany( @@ -1021,10 +1065,24 @@ def write_records(conn: sqlite3.Connection, rec: SessionRecords, now_ms: int) -> " schema_version=excluded.schema_version, " " parser_version=excluded.parser_version", ( - sf, m["folder"], m["is_subagent"], m["parent_session"], m["subagent_label"], - m["started_at"], m["title"], m["cwd"], m["session_uuid"], m["version"], - rec.file_mtime, rec.file_size, rec.final_offset, rec.final_line_count, - now_ms, TOKENIZER_NAME, SCHEMA_VERSION, EDIT_PARSER_VERSION, + sf, + m["folder"], + m["is_subagent"], + m["parent_session"], + m["subagent_label"], + m["started_at"], + m["title"], + m["cwd"], + m["session_uuid"], + m["version"], + rec.file_mtime, + rec.file_size, + rec.final_offset, + rec.final_line_count, + now_ms, + TOKENIZER_NAME, + SCHEMA_VERSION, + EDIT_PARSER_VERSION, ), ) cur.execute("COMMIT") @@ -1036,6 +1094,7 @@ def write_records(conn: sqlite3.Connection, rec: SessionRecords, now_ms: int) -> # --------------------------------------------------------------------------- # # Driver + def discover_sessions(root: Path, limit: int | None) -> list[Path]: if not root.exists(): return [] @@ -1077,18 +1136,27 @@ def decide_action( def main() -> int: ap = argparse.ArgumentParser() ap.add_argument("--workers", type=int, default=min(16, (os.cpu_count() or 4) * 2)) - ap.add_argument("--limit", type=int, default=0, - help="only sync the N most-recent files (0 = all)") - ap.add_argument("--full", action="store_true", - help="ignore stored state, re-ingest every file from scratch") + ap.add_argument( + "--limit", + type=int, + default=0, + help="only sync the N most-recent files (0 = all)", + ) + ap.add_argument( + "--full", + action="store_true", + help="ignore stored state, re-ingest every file from scratch", + ) ap.add_argument("--root", default=str(SESSIONS_ROOT)) args = ap.parse_args() root = Path(args.root).expanduser() print(f"-> sessions root: {root}", file=sys.stderr) print(f"-> db: {DB_PATH}", file=sys.stderr) - print(f"-> parser_version={EDIT_PARSER_VERSION} schema_version={SCHEMA_VERSION}", - file=sys.stderr) + print( + f"-> parser_version={EDIT_PARSER_VERSION} schema_version={SCHEMA_VERSION}", + file=sys.stderr, + ) conn = open_db() state = existing_state(conn) @@ -1123,14 +1191,18 @@ def main() -> int: try: write_records(conn, rec, now_ms) except Exception as e: - print(f"!! write failed for {rec.session_meta['session_file']}: {e}", - file=sys.stderr) + print( + f"!! write failed for {rec.session_meta['session_file']}: {e}", + file=sys.stderr, + ) n += 1 now = time.monotonic() if now - last_log >= 1.0: rate = n / max(now - t0, 1e-6) - print(f" wrote {n}/{len(work)} files ({rate:.1f} files/s)", - file=sys.stderr) + print( + f" wrote {n}/{len(work)} files ({rate:.1f} files/s)", + file=sys.stderr, + ) last_log = now rate = n / max(time.monotonic() - t0, 1e-6) print(f"-> wrote {n} files total ({rate:.1f} files/s)", file=sys.stderr) diff --git a/scripts/tool_io.py b/scripts/tool_io.py index 332f5beb7..ef10895cd 100755 --- a/scripts/tool_io.py +++ b/scripts/tool_io.py @@ -134,7 +134,6 @@ class ReservoirSample[T]: self.items[index] = item - def list_recent_session_files(config: ToolIOConfig) -> list[Path]: min_mtime = time.time() - config.since_days * 24 * 60 * 60 candidates: list[tuple[float, Path]] = [] @@ -150,7 +149,6 @@ def list_recent_session_files(config: ToolIOConfig) -> list[Path]: return [entry[1] for entry in candidates[: config.max_files]] - def iter_tool_invocations( tool_names: str | Iterable[str], config: ToolIOConfig | None = None, @@ -183,7 +181,11 @@ def iter_tool_invocations( continue tool_name = _as_string(payload.get("name")) tool_call_id = _as_string(payload.get("id")) - if tool_name is None or tool_call_id is None or tool_name not in wanted: + if ( + tool_name is None + or tool_call_id is None + or tool_name not in wanted + ): continue arguments = _as_record(payload.get("arguments")) or {} pending[tool_call_id] = ToolCall( @@ -230,43 +232,39 @@ def iter_tool_invocations( return - def iter_results(stream: Iterable[ToolInvocation]) -> Iterator[ToolInvocation]: for invocation in stream: if invocation.has_result: yield invocation - def iter_failed(stream: Iterable[ToolInvocation]) -> Iterator[ToolInvocation]: for invocation in stream: if invocation.is_error: yield invocation - def iter_successful(stream: Iterable[ToolInvocation]) -> Iterator[ToolInvocation]: for invocation in stream: if invocation.has_result and not invocation.is_error: yield invocation - def iter_with_diff(stream: Iterable[ToolInvocation]) -> Iterator[ToolInvocation]: for invocation in stream: if invocation.diff: yield invocation - -def iter_paths(stream: Iterable[ToolInvocation], *paths: str) -> Iterator[ToolInvocation]: +def iter_paths( + stream: Iterable[ToolInvocation], *paths: str +) -> Iterator[ToolInvocation]: wanted = set(paths) for invocation in stream: if invocation.path_hint in wanted: yield invocation - def take(stream: Iterable[ToolInvocation], limit: int) -> Iterator[ToolInvocation]: if limit <= 0: return @@ -278,15 +276,15 @@ def take(stream: Iterable[ToolInvocation], limit: int) -> Iterator[ToolInvocatio remaining -= 1 - -def sample_reservoir[T](stream: Iterable[T], size: int, seed: int | None = None) -> list[T]: +def sample_reservoir[T]( + stream: Iterable[T], size: int, seed: int | None = None +) -> list[T]: sample: ReservoirSample[T] = ReservoirSample(size=size, rng=random.Random(seed)) for item in stream: sample.add(item) return sample.items - def extract_result_text(message: dict[str, Any] | None) -> str: if message is None: return "" @@ -305,7 +303,6 @@ def extract_result_text(message: dict[str, Any] | None) -> str: return "" - def extract_path(arguments: dict[str, Any]) -> str: for key in ("path", "file", "move"): value = arguments.get(key) @@ -314,7 +311,6 @@ def extract_path(arguments: dict[str, Any]) -> str: return "" - def _iter_session_entries(session_file: Path) -> Iterator[dict[str, Any]]: with session_file.open("r", encoding="utf-8") as handle: for line in handle: @@ -330,7 +326,6 @@ def _iter_session_entries(session_file: Path) -> Iterator[dict[str, Any]]: yield payload - def _extract_thinking(content: list[Any]) -> str | None: for item in content: payload = _as_record(item) @@ -344,7 +339,6 @@ def _extract_thinking(content: list[Any]) -> str | None: return None - def resolve_tool_names(*names_or_groups: str) -> tuple[str, ...]: ordered: list[str] = [] seen: set[str] = set() @@ -367,20 +361,17 @@ def _normalize_tool_names(tool_names: str | Iterable[str]) -> set[str]: return set(ordered) - def _event_weight(invocation: ToolInvocation, limit_mode: LimitMode) -> int: if limit_mode == "calls": return 1 return 2 if invocation.has_result else 1 - def _as_record(value: Any) -> dict[str, Any] | None: if not isinstance(value, dict): return None return value - def _as_string(value: Any) -> str | None: return value if isinstance(value, str) else None