loupe / read_runtime.py
anxuanng's picture
Loupe research preview: audit code, read weights, benchmark reports
84c3565 verified
Raw History Blame Contribute Delete
8.39 kB
"""Everything needed to run Loupe's trained read, in one file with no research-code imports.
The read takes two named numbers (say a heading's and its text's font size), a requested band for their
ratio, and a question; it answers in range, too low or too high with a probability. The weights are
read_v2/read.safetensors (about 121 thousand parameters). The question and option sentences were encoded once
by a text encoder and are stored in text_cache.pt, so no text model runs at audit time.
The model and packing code are the ones the read was trained with (read_model.py, train_read.py,
website_design_model_v10.py, website_brief_inputs.py), copied here without the training parts.
"""
from pathlib import Path
import torch
from torch import Tensor, nn
PACKAGE = Path(__file__).resolve().parent
TEXT = torch.load(PACKAGE / "text_cache.pt", map_location="cpu", weights_only=True)
RELATION_SLOTS = 8
JUDGE_OPTIONS = ("yes", "no")
ADJUST_OPTIONS = ("increase", "decrease", "keep")
FIXED_KINDS = {"typography": ("heading_font_px", "body_font_px"), "leading": ("body_line_height_px", "body_font_px"),
"grouping": ("heading_body_gap_px", "next_group_gap_px")}
def kind_bytes(kind):
"""A quantity name as 64 byte slots, so two names match only when they are the same string."""
encoded_kind = kind.encode("utf-8")
if not encoded_kind or len(encoded_kind) > 64:
raise ValueError("Brief quantity kind requires 1 to 64 UTF-8 bytes")
packed = torch.zeros(64, dtype=torch.long)
packed[:len(encoded_kind)] = torch.tensor([byte + 1 for byte in encoded_kind])
return packed
def select(kind: Tensor, relation_kind_bytes: Tensor, log_values: Tensor, relation_mask: Tensor, required: Tensor) -> Tensor:
"""The log value of the one supplied quantity whose name equals `kind`."""
matches = (relation_kind_bytes == kind.unsqueeze(1)).all(dim=-1) & relation_mask
if not bool((matches.sum(dim=-1) == 1)[required].all()):
raise ValueError("Each brief quantity must match exactly one supplied relation")
return (matches * log_values).sum(dim=-1)
class MarginBriefRead(nn.Module):
"""Exact relation lookup, margins to both bounds by subtraction, learned scoring of each option."""
def __init__(self, text_width: int, hidden_width: int = 64) -> None:
super().__init__()
self.numeric = nn.Sequential(
nn.Linear(6, hidden_width), nn.GELU(), nn.Linear(hidden_width, hidden_width), nn.GELU(),
nn.Linear(hidden_width, hidden_width),
)
self.text_norm = nn.LayerNorm(text_width)
self.question_text = nn.Linear(text_width, hidden_width)
self.option_text = nn.Linear(text_width, hidden_width)
self.score = nn.Sequential(
nn.Linear(hidden_width * 3, hidden_width), nn.GELU(), nn.Linear(hidden_width, 1)
)
def forward(
self, question_text: Tensor, option_text: Tensor, relation_kind_bytes: Tensor,
relation_values: Tensor, relation_mask: Tensor, brief_numerator_kind: Tensor,
brief_denominator_kind: Tensor, brief_low: Tensor, brief_high: Tensor,
brief_mask: Tensor | None = None,
) -> Tensor:
relation_mask = relation_mask.bool()
required = torch.ones_like(brief_low, dtype=torch.bool) if brief_mask is None else brief_mask.bool()
raw_values = (relation_values[..., 4].double() * 100).clamp_min(1e-6)
log_values = torch.where(relation_mask, raw_values.log(), torch.zeros_like(raw_values))
numerator = select(brief_numerator_kind, relation_kind_bytes, log_values, relation_mask, required)
denominator = select(brief_denominator_kind, relation_kind_bytes, log_values, relation_mask, required)
log_ratio = numerator - denominator
# Positive when the ratio is above the low bound, and when it is below the high bound.
above_low = (log_ratio - brief_low.double().log()).float()
below_high = (brief_high.double().log() - log_ratio).float()
numbers = torch.stack([above_low, below_high, torch.tanh(above_low * 40), torch.tanh(below_high * 40),
torch.tanh(above_low * 400), torch.tanh(below_high * 400)], dim=-1)
brief = self.numeric(numbers)[:, None, None, :]
candidates = self.question_text(self.text_norm(question_text)).unsqueeze(2)
candidates = candidates + self.option_text(self.text_norm(option_text))
brief = brief.expand_as(candidates)
return self.score(torch.cat([candidates, brief, candidates * brief], dim=-1)).squeeze(-1)
def embedding(sentence):
return TEXT["text"][TEXT["strings"].index(sentence)]
def question_sets():
"""(judge sentence, adjust sentence, fixed kinds or None): the generic wording and the three original ones."""
sets = [(TEXT["generic"]["judge"], TEXT["generic"]["adjust"], None)]
for criterion, subject in TEXT["criteria"].items():
sets.append((f"Does the measured {subject} fall within the requested experimental policy band?",
f"Choose how to adjust the {subject} to satisfy the requested policy band.", FIXED_KINDS[criterion]))
return sets
QUESTION_SETS = question_sets()
OPTION_FEATURES = {name: embedding(sentence) for name, sentence in TEXT["options"].items()}
def uniform(generator, low, high):
return float(torch.empty(1).uniform_(low, high, generator=generator))
def policy(ratio, low, high):
if ratio < low:
return "no", "increase"
if ratio > high:
return "no", "decrease"
return "yes", "keep"
def pack(items, generator):
"""Tensors for the read from a list of {set, kinds, relations, low, high, ratio}; orders are shuffled."""
count = len(items)
tensors = {
"question_text": torch.zeros(count, 2, 768), "option_text": torch.zeros(count, 2, 3, 768),
"option_mask": torch.zeros(count, 2, 3, dtype=torch.bool), "targets": torch.zeros(count, 2, 3),
"relation_kind_bytes": torch.zeros(count, RELATION_SLOTS, 64, dtype=torch.long),
"relation_values": torch.zeros(count, RELATION_SLOTS, 5),
"relation_mask": torch.zeros(count, RELATION_SLOTS, dtype=torch.bool),
"brief_numerator_kind": torch.zeros(count, 64, dtype=torch.long),
"brief_denominator_kind": torch.zeros(count, 64, dtype=torch.long),
"brief_low": torch.zeros(count), "brief_high": torch.zeros(count),
}
for row, item in enumerate(items):
judge_sentence, adjust_sentence, _ = QUESTION_SETS[item["set"]]
judge_answer, adjust_answer = policy(item["ratio"], item["low"], item["high"])
questions = [(judge_sentence, JUDGE_OPTIONS, judge_answer), (adjust_sentence, ADJUST_OPTIONS, adjust_answer)]
if uniform(generator, 0, 1) < 0.5:
questions.reverse()
for column, (sentence, options, answer) in enumerate(questions):
tensors["question_text"][row, column] = embedding(sentence)
order = torch.randperm(len(options), generator=generator).tolist()
for position, option_index in enumerate(order):
name = options[option_index]
tensors["option_text"][row, column, position] = OPTION_FEATURES[name]
tensors["option_mask"][row, column, position] = True
tensors["targets"][row, column, position] = float(name == answer)
relation_order = torch.randperm(len(item["relations"]), generator=generator).tolist()
for position, index in enumerate(relation_order):
kind, value = item["relations"][index]
tensors["relation_kind_bytes"][row, position] = kind_bytes(kind)
tensors["relation_values"][row, position, 4] = value / 100
tensors["relation_mask"][row, position] = True
tensors["brief_numerator_kind"][row] = kind_bytes(item["kinds"][0])
tensors["brief_denominator_kind"][row] = kind_bytes(item["kinds"][1])
tensors["brief_low"][row] = item["low"]
tensors["brief_high"][row] = item["high"]
return tensors
def logits_of(read, tensors):
names = ("question_text", "option_text", "relation_kind_bytes", "relation_values", "relation_mask",
"brief_numerator_kind", "brief_denominator_kind", "brief_low", "brief_high")
return read(**{name: tensors[name] for name in names}).masked_fill(~tensors["option_mask"], -1e9)