marcodsn commited on
Commit
ec158b2
·
verified ·
1 Parent(s): 75b5f26

alpha_sys_1.py: shared-prefix inference for multi-question requests

Browse files
Files changed (1) hide show
  1. alpha_sys_1.py +59 -10
alpha_sys_1.py CHANGED
@@ -12,6 +12,8 @@ on them and reads the answer distribution from one forward pass.
12
  m.system_one({"state": ..., "images": [...], "questions": {"q1": {...}, "q2": {...}}})
13
  # -> {"model": ..., "answers": {"q1": {...}, "q2": {...}}} (TypeSafe's System One shape)
14
 
 
 
15
  Question types: choice (criteria = {option: description or None} or a list of options, up to
16
  26), noul (a statement; criteria = {"true": ..., "false": ...} optional), score (criteria = the
17
  levels, lowest first; the score is the expected level index). Images: a PIL image, a path, or
@@ -88,9 +90,13 @@ def confidence(p: list[float]) -> float:
88
 
89
 
90
  class SystemOne:
91
- def __init__(self, repo: str, revision: str | None = None, device: str | None = None, dtype=torch.bfloat16, temperature: float = 1.0):
92
- """temperature: the label logits are divided by it (1.0 = the model as released; see fit_temperature)."""
93
- self.repo, self.revision, self.temperature = repo, revision, temperature
 
 
 
 
94
  self.processor = AutoProcessor.from_pretrained(repo, revision=revision)
95
  self.processor.tokenizer.padding_side = "left"
96
  self.model = AutoModelForImageTextToText.from_pretrained(repo, revision=revision, dtype=dtype)
@@ -105,25 +111,68 @@ class SystemOne:
105
  self._ids[label] = ids[0]
106
  return self._ids[label]
107
 
108
- @torch.inference_mode()
109
- def distributions(self, items: list[tuple[Any, dict, list | None]]) -> list[list[float]]:
110
- """items: (state, question, images or None) -> probabilities in the answer-space order."""
111
  msgs, labels_per = [], []
112
  for state, q, images in items:
113
  text, labels, _ = render(state, q)
114
  content = [{"type": "image", "image": load_image(im)} for im in (images or [])] + [{"type": "text", "text": text}]
115
  msgs.append([{"role": "user", "content": content}])
116
  labels_per.append(labels)
117
- inputs = self.processor.apply_chat_template(
118
- msgs, add_generation_prompt=True, tokenize=True, return_dict=True,
119
- processor_kwargs={"return_tensors": "pt", "padding": True}).to(self.device)
120
- logits = self.model(**inputs, logits_to_keep=1).logits[:, -1].float()
121
  out = []
122
  for i, labels in enumerate(labels_per):
123
  ids = torch.tensor([self._label_id(lab) for lab in labels], device=logits.device)
124
  out.append(torch.softmax(logits[i, ids] / self.temperature, -1).tolist())
125
  return out
126
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
127
  def answer(self, q: dict, p: list[float]) -> dict:
128
  _, _, keys = render(None, q)
129
  if q["type"] == "choice":
 
12
  m.system_one({"state": ..., "images": [...], "questions": {"q1": {...}, "q2": {...}}})
13
  # -> {"model": ..., "answers": {"q1": {...}, "q2": {...}}} (TypeSafe's System One shape)
14
 
15
+ The questions of one request share their images and state, so that prefix is computed once
16
+ and only the questions run against its cache (`shared_prefix=False` turns this off).
17
  Question types: choice (criteria = {option: description or None} or a list of options, up to
18
  26), noul (a statement; criteria = {"true": ..., "false": ...} optional), score (criteria = the
19
  levels, lowest first; the score is the expected level index). Images: a PIL image, a path, or
 
90
 
91
 
92
  class SystemOne:
93
+ def __init__(self, repo: str, revision: str | None = None, device: str | None = None, dtype=torch.bfloat16, temperature: float = 1.0,
94
+ shared_prefix: bool = True):
95
+ """temperature: the label logits are divided by it (1.0 = the model as released; see fit_temperature).
96
+ shared_prefix: when several questions share their images and state, run that prefix once and
97
+ only the questions against its cache (same probabilities up to bfloat16 noise, several times
98
+ faster with an image). False runs every question as its own full sequence."""
99
+ self.repo, self.revision, self.temperature, self.shared_prefix = repo, revision, temperature, shared_prefix
100
  self.processor = AutoProcessor.from_pretrained(repo, revision=revision)
101
  self.processor.tokenizer.padding_side = "left"
102
  self.model = AutoModelForImageTextToText.from_pretrained(repo, revision=revision, dtype=dtype)
 
111
  self._ids[label] = ids[0]
112
  return self._ids[label]
113
 
114
+ def _messages(self, items: list[tuple[Any, dict, list | None]]) -> tuple[list, list[list[str]]]:
 
 
115
  msgs, labels_per = [], []
116
  for state, q, images in items:
117
  text, labels, _ = render(state, q)
118
  content = [{"type": "image", "image": load_image(im)} for im in (images or [])] + [{"type": "text", "text": text}]
119
  msgs.append([{"role": "user", "content": content}])
120
  labels_per.append(labels)
121
+ return msgs, labels_per
122
+
123
+ def _probs(self, logits: torch.Tensor, labels_per: list[list[str]]) -> list[list[float]]:
 
124
  out = []
125
  for i, labels in enumerate(labels_per):
126
  ids = torch.tensor([self._label_id(lab) for lab in labels], device=logits.device)
127
  out.append(torch.softmax(logits[i, ids] / self.temperature, -1).tolist())
128
  return out
129
 
130
+ @torch.inference_mode()
131
+ def distributions(self, items: list[tuple[Any, dict, list | None]]) -> list[list[float]]:
132
+ """items: (state, question, images or None) -> probabilities in the answer-space order.
133
+ Every item is its own sequence, read at its last position; items that share images and
134
+ state share the prefix's computation when `shared_prefix` is on."""
135
+ msgs, labels_per = self._messages(items)
136
+ if self.shared_prefix and len(items) > 1 and all(it[0] == items[0][0] and (it[2] or None) == (items[0][2] or None) for it in items):
137
+ logits = self._shared_prefix_logits(msgs)
138
+ if logits is not None:
139
+ return self._probs(logits, labels_per)
140
+ inputs = self.processor.apply_chat_template(
141
+ msgs, add_generation_prompt=True, tokenize=True, return_dict=True,
142
+ processor_kwargs={"return_tensors": "pt", "padding": True}).to(self.device)
143
+ return self._probs(self.model(**inputs, logits_to_keep=1).logits[:, -1].float(), labels_per)
144
+
145
+ def _shared_prefix_logits(self, msgs: list) -> torch.Tensor | None:
146
+ """One pass over the longest common token prefix (images, state), then the question suffixes,
147
+ right-padded, against that cache repeated across the batch. None when there is too little to share."""
148
+ encs = [self.processor.apply_chat_template([m], add_generation_prompt=True, tokenize=True, return_dict=True,
149
+ processor_kwargs={"return_tensors": "pt"}) for m in msgs]
150
+ ids = [e["input_ids"][0] for e in encs]
151
+ L = min(len(x) for x in ids) - 1 # at least one token per suffix
152
+ for x in ids[1:]:
153
+ diff = (x[:L] != ids[0][:L]).nonzero()
154
+ if len(diff):
155
+ L = min(L, int(diff[0]))
156
+ image_id = getattr(self.model.config, "image_token_id", None)
157
+ if L < 64 or (image_id is not None and any((x[L:] == image_id).any() for x in ids)):
158
+ return None
159
+ n = len(ids)
160
+ image_kw = {k: v.to(self.device) for k, v in encs[0].items() if k in ("pixel_values", "spatial_shapes", "pixel_attention_mask")}
161
+ cache = self.model(input_ids=ids[0][:L][None].to(self.device), **image_kw, use_cache=True).past_key_values
162
+ cache.reorder_cache(torch.zeros(n, dtype=torch.long, device=self.device))
163
+ sufs = [x[L:] for x in ids]
164
+ lens = torch.tensor([len(s) for s in sufs])
165
+ width = int(lens.max())
166
+ suffix = torch.full((n, width), self.processor.tokenizer.pad_token_id, dtype=torch.long)
167
+ attention = torch.zeros((n, L + width), dtype=torch.long)
168
+ attention[:, :L] = 1
169
+ for i, s in enumerate(sufs):
170
+ suffix[i, : len(s)] = s
171
+ attention[i, L : L + len(s)] = 1
172
+ out = self.model(input_ids=suffix.to(self.device), attention_mask=attention.to(self.device), past_key_values=cache,
173
+ cache_position=torch.arange(L, L + width, device=self.device), use_cache=True)
174
+ return out.logits[torch.arange(n), (lens - 1).to(self.device)].float()
175
+
176
  def answer(self, q: dict, p: list[float]) -> dict:
177
  _, _, keys = render(None, q)
178
  if q["type"] == "choice":