"""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)