File size: 8,393 Bytes
84c3565 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 | """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)
|