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)