larkooo commited on
Commit
f092803
·
verified ·
1 Parent(s): 53e24ca

Add simultaneous streaming image and video comparisons

Browse files
.gitattributes CHANGED
@@ -36,3 +36,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
36
  docs/assets/demo-poster.jpg filter=lfs diff=lfs merge=lfs -text
37
  docs/assets/demo.mp4 filter=lfs diff=lfs merge=lfs -text
38
  tokenizer.json filter=lfs diff=lfs merge=lfs -text
 
 
36
  docs/assets/demo-poster.jpg filter=lfs diff=lfs merge=lfs -text
37
  docs/assets/demo.mp4 filter=lfs diff=lfs merge=lfs -text
38
  tokenizer.json filter=lfs diff=lfs merge=lfs -text
39
+ gemma_rlcd/static/sample-street.jpg filter=lfs diff=lfs merge=lfs -text
MANIFEST.in CHANGED
@@ -3,6 +3,6 @@ recursive-include docs *.md *.txt *.jpg *.mp4 *.json
3
  recursive-include examples *.json *.jsonl *.txt
4
  recursive-include reports *.md *.json
5
  recursive-include scripts *.py
6
- recursive-include tests *.py
7
- recursive-include gemma_rlcd/static *.html *.css *.js *.json
8
  global-exclude .DS_Store *.py[cod] *.safetensors *.gguf
 
3
  recursive-include examples *.json *.jsonl *.txt
4
  recursive-include reports *.md *.json
5
  recursive-include scripts *.py
6
+ recursive-include tests *.py *.cjs
7
+ recursive-include gemma_rlcd/static *.html *.css *.js *.json *.jpg
8
  global-exclude .DS_Store *.py[cod] *.safetensors *.gguf
README.md CHANGED
@@ -52,6 +52,10 @@ Choose **Run all fields** for answers and distributions, or **Compare with Gemma
52
 
53
  The download is approximately 3.6 GB and includes the image and audio encoders, tokenizer, processor, and chat template. Inference runs locally on Apple Silicon through MLX.
54
 
 
 
 
 
55
  ## How it works
56
 
57
  ```mermaid
@@ -130,7 +134,7 @@ A ratio above 1 favors parallel scoring. The support example matches all 28 outp
130
 
131
  The playground accepts up to eight images, one audio source, and one video, with 200 MB of uploads per request. Audio and videos with sound support up to 30 seconds; silent video supports up to 60 seconds. Video targets one frame per second, capped at 32 frames, so brief events can fall between samples.
132
 
133
- Requests support up to 32 named fields, 64 primitive decisions, and 8,192 processed input tokens. Oversized inputs return an error rather than being truncated. Multi-picture phone JPEGs use the full-resolution primary photograph. Text and field definitions are saved in browser local storage; uploaded media is not retained.
134
 
135
  ## Development
136
 
 
52
 
53
  The download is approximately 3.6 GB and includes the image and audio encoders, tokenizer, processor, and chat template. Inference runs locally on Apple Silicon through MLX.
54
 
55
+ ### Live visual demo
56
+
57
+ Open **http://127.0.0.1:8787/demo** for 32, 64, or 128 checks over an image or video. Use the included street photo or upload your own media. The parallel scorer streams completed field batches; normal Gemma streams its generated JSON. Live clocks, per-check probabilities, answer differences, and downloadable events make the comparison inspectable. Both paths start together and stream side by side, sharing the resident weights with separate processors and KV caches. Timings measure concurrent completion on one GPU, including resource contention. The playground’s ordinary comparison remains sequential for isolated timings.
58
+
59
  ## How it works
60
 
61
  ```mermaid
 
134
 
135
  The playground accepts up to eight images, one audio source, and one video, with 200 MB of uploads per request. Audio and videos with sound support up to 30 seconds; silent video supports up to 60 seconds. Video targets one frame per second, capped at 32 frames, so brief events can fall between samples.
136
 
137
+ Requests support up to 32 named fields, 128 primitive decisions, and 8,192 processed input tokens. Oversized inputs return an error rather than being truncated. Multi-picture phone JPEGs use the full-resolution primary photograph. Text and field definitions are saved in browser local storage; uploaded media is not retained.
138
 
139
  ## Development
140
 
THIRD_PARTY_NOTICES.md CHANGED
@@ -45,6 +45,10 @@ The project's MIT license applies to project code and synthetic examples, not th
45
 
46
  The demo video and poster use IBM Plex Sans, licensed under the SIL Open Font License. The license is retained in [docs/assets/Plex-OFL.txt](docs/assets/Plex-OFL.txt).
47
 
 
 
 
 
48
  ## Bundled checkpoint
49
 
50
  The complete Gemma 4 E2B MLX checkpoint is included unchanged from the revision in [checkpoint provenance](checkpoint-provenance.json). Model licensing and attribution are retained in [MODEL_LICENSE](MODEL_LICENSE) and [NOTICE](NOTICE). The original conversion card is preserved in [docs/upstream-mlx-model-card.md](docs/upstream-mlx-model-card.md).
 
45
 
46
  The demo video and poster use IBM Plex Sans, licensed under the SIL Open Font License. The license is retained in [docs/assets/Plex-OFL.txt](docs/assets/Plex-OFL.txt).
47
 
48
+ ## Visual demo photograph
49
+
50
+ `gemma_rlcd/static/sample-street.jpg` is “Times Square (New York City)” by ISO Legacy, from [Wikimedia Commons](https://commons.wikimedia.org/wiki/File:Times_Square_(New_York_City).jpg), dedicated under [CC0 1.0](https://creativecommons.org/publicdomain/zero/1.0/). The original photograph is included without modification.
51
+
52
  ## Bundled checkpoint
53
 
54
  The complete Gemma 4 E2B MLX checkpoint is included unchanged from the revision in [checkpoint provenance](checkpoint-provenance.json). Model licensing and attribution are retained in [MODEL_LICENSE](MODEL_LICENSE) and [NOTICE](NOTICE). The original conversion card is preserved in [docs/upstream-mlx-model-card.md](docs/upstream-mlx-model-card.md).
checkpoint-provenance.json CHANGED
@@ -1,7 +1,7 @@
1
  {
2
  "repository": "larkooo/gemma-e2b-rlcd",
3
  "runtime_repository": "https://github.com/Larkooo/gemma-e2b-rlcd",
4
- "runtime_commit": "2eb9b959e989b08eb140578282336223243a6084",
5
  "checkpoint_repository": "mlx-community/gemma-4-e2b-it-4bit",
6
  "checkpoint_revision": "238767527555cb75a05732a84dff5d6ba0dd6809",
7
  "checkpoint_modified": false,
 
1
  {
2
  "repository": "larkooo/gemma-e2b-rlcd",
3
  "runtime_repository": "https://github.com/Larkooo/gemma-e2b-rlcd",
4
+ "runtime_commit": "e78b79605ac4e37b776727935e2de7dc8f259a81",
5
  "checkpoint_repository": "mlx-community/gemma-4-e2b-it-4bit",
6
  "checkpoint_revision": "238767527555cb75a05732a84dff5d6ba0dd6809",
7
  "checkpoint_modified": false,
docs/architecture.md CHANGED
@@ -53,7 +53,7 @@ Float32 reduced the probability drift observed when changing execution shapes un
53
  - One silent video up to 60 seconds, sampled at a target of 1 fps with a 32-frame processor cap.
54
  - A video's soundtrack is included automatically. Videos with audio have a 30-second limit; a second audio source returns an error.
55
  - Up to 8,192 processed input tokens by default. Inputs exceeding the limit are rejected without truncation.
56
- - Up to 32 named fields and 64 primitive decisions in the playground.
57
  - Questions and descriptions are strings; the Python contracts use name-to-description mappings.
58
 
59
  Video sampling can miss brief events. The limits describe the current serving configuration rather than the checkpoint's maximum context capacity.
 
53
  - One silent video up to 60 seconds, sampled at a target of 1 fps with a 32-frame processor cap.
54
  - A video's soundtrack is included automatically. Videos with audio have a 30-second limit; a second audio source returns an error.
55
  - Up to 8,192 processed input tokens by default. Inputs exceeding the limit are rejected without truncation.
56
+ - Up to 32 named fields and 128 primitive decisions in the playground.
57
  - Questions and descriptions are strings; the Python contracts use name-to-description mappings.
58
 
59
  Video sampling can miss brief events. The limits describe the current serving configuration rather than the checkpoint's maximum context capacity.
gemma_rlcd/cached_backend.py CHANGED
@@ -149,7 +149,11 @@ class CachedMLXBackend(MLXBackend):
149
  return scores
150
 
151
  def branches(
152
- self, prepared: PreparedState, prefix_cache, requests: Sequence[ScoringRequest]
 
 
 
 
153
  ) -> list[TokenScores]:
154
  mx = self.mx
155
  scores = []
@@ -188,9 +192,10 @@ class CachedMLXBackend(MLXBackend):
188
  :, 0, :
189
  ].astype(mx.float32)
190
  mx.eval(logits)
191
- scores.extend(
192
- self._extract(logits, batch, [prepared.prefix_tokens + n for n in lengths])
193
- )
 
194
  batch_sizes.append(len(batch))
195
  tail_lengths.append(tail_length)
196
  self.last_stats["branch_batch_sizes"] = batch_sizes
 
149
  return scores
150
 
151
  def branches(
152
+ self,
153
+ prepared: PreparedState,
154
+ prefix_cache,
155
+ requests: Sequence[ScoringRequest],
156
+ on_batch=None,
157
  ) -> list[TokenScores]:
158
  mx = self.mx
159
  scores = []
 
192
  :, 0, :
193
  ].astype(mx.float32)
194
  mx.eval(logits)
195
+ completed = self._extract(logits, batch, [prepared.prefix_tokens + n for n in lengths])
196
+ scores.extend(completed)
197
+ if on_batch is not None:
198
+ on_batch(start, completed)
199
  batch_sizes.append(len(batch))
200
  tail_lengths.append(tail_length)
201
  self.last_stats["branch_batch_sizes"] = batch_sizes
gemma_rlcd/comparison.py CHANGED
@@ -2,6 +2,10 @@
2
 
3
  import json
4
  import time
 
 
 
 
5
 
6
  from .core import Choice, DecisionEngine, Independent, Noul, Score, State
7
  from .mlx_backend import audio_paths
@@ -140,17 +144,14 @@ def prepare_generation(backend, state: State, questions: dict) -> tuple[str, dic
140
  return prompt, inputs
141
 
142
 
143
- def generate_answers(backend, state: State, questions: dict) -> dict:
144
- from mlx_vlm import generate
145
 
146
  started = time.perf_counter()
147
  prompt, inputs = prepare_generation(backend, state, questions)
148
  budget = output_budget(backend.tokenizer, questions)
149
  prepared = time.perf_counter()
150
- generated = generate(
151
- backend.model,
152
- backend.processor,
153
- prompt,
154
  **inputs,
155
  max_tokens=budget,
156
  temperature=0,
@@ -158,19 +159,32 @@ def generate_answers(backend, state: State, questions: dict) -> dict:
158
  logits_to_keep=1,
159
  verbose=False,
160
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
161
  backend.mx.synchronize()
162
  generated_at = time.perf_counter()
163
  error = None
164
  answers = None
165
  try:
166
- answers = parse_generated(generated.text, questions)
167
  if generated.finish_reason != "stop":
168
  error = "Generation reached its token limit without an end-of-answer token"
169
  except ValueError as exc:
170
  error = str(exc)
171
  return {
172
  "answers": answers,
173
- "raw_text": generated.text,
174
  "valid": error is None,
175
  "error": error,
176
  "finish_reason": generated.finish_reason,
@@ -197,30 +211,132 @@ def discrete_answers(answers: dict) -> dict:
197
  return values
198
 
199
 
200
- def compare(backend, state: State, questions: dict, media_seconds: float) -> tuple[dict, dict]:
201
- def evaluate(method):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
202
  # No cross-run KV or vision cache. Start each path with completed GPU
203
  # work and a cleared allocator cache; weights remain resident.
204
- if hasattr(backend, "mx"):
205
- backend.mx.synchronize()
206
- backend.mx.clear_cache()
207
- started = time.perf_counter()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
208
  if method == "parallel":
209
- output = DecisionEngine(backend).system_one(state, questions)
210
- if hasattr(backend, "mx"):
211
- backend.mx.synchronize()
 
 
 
 
 
212
  output.update(
213
  inference_seconds=time.perf_counter() - started,
214
- execution=dict(backend.last_stats),
215
  valid=True,
216
  )
217
  else:
218
- output = generate_answers(backend, state, questions)
 
 
 
 
 
 
219
  output["total_seconds"] = media_seconds + output["inference_seconds"]
 
 
 
 
 
 
 
 
 
220
  return output
221
 
222
- parallel = evaluate("parallel")
223
- normal = evaluate("normal")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
224
  decisions = discrete_answers(parallel["answers"])
225
  return parallel, {
226
  "seconds": {"parallel": parallel["total_seconds"], "normal": normal["total_seconds"]},
@@ -234,8 +350,14 @@ def compare(backend, state: State, questions: dict, media_seconds: float) -> tup
234
  for name in questions
235
  },
236
  "methodology": {
 
237
  "model": "Same resident Gemma 4 E2B 4-bit weights; float32 compute; all 35 layers",
238
- "timing": "One run per path, parallel scorer first, normal generation second. Includes input preparation and inference; shared upload decoding added equally to each path. Excludes model loading, upload transfer, and allocator reset. No warm-up runs; first-use effects and run order can affect this observation.",
 
 
 
 
 
239
  "cache": "Fresh input KV and media features for every run; ordinary output-token KV caching remains enabled for normal generation.",
240
  "answers": "Normal Gemma generates one JSON object for all fields, greedily, without thinking. Compare choice names, most likely grade levels, and booleans at a 50% threshold. The scorer also returns probability distributions and expected grades. Agreement is not an accuracy measurement.",
241
  },
 
2
 
3
  import json
4
  import time
5
+ from concurrent.futures import ThreadPoolExecutor
6
+ from contextlib import nullcontext
7
+ from copy import copy, deepcopy
8
+ from threading import Barrier
9
 
10
  from .core import Choice, DecisionEngine, Independent, Noul, Score, State
11
  from .mlx_backend import audio_paths
 
144
  return prompt, inputs
145
 
146
 
147
+ def generate_answers(backend, state: State, questions: dict, on_token=None) -> dict:
148
+ from mlx_vlm import generate, stream_generate
149
 
150
  started = time.perf_counter()
151
  prompt, inputs = prepare_generation(backend, state, questions)
152
  budget = output_budget(backend.tokenizer, questions)
153
  prepared = time.perf_counter()
154
+ options = dict(
 
 
 
155
  **inputs,
156
  max_tokens=budget,
157
  temperature=0,
 
159
  logits_to_keep=1,
160
  verbose=False,
161
  )
162
+ if on_token is None:
163
+ generated = generate(backend.model, backend.processor, prompt, **options)
164
+ text = generated.text
165
+ else:
166
+ parts = []
167
+ generated = None
168
+ for chunk in stream_generate(backend.model, backend.processor, prompt, **options):
169
+ generated = chunk
170
+ parts.append(chunk.text)
171
+ on_token(chunk.text, chunk.generation_tokens)
172
+ if generated is None:
173
+ raise RuntimeError("Gemma returned no generation result")
174
+ text = "".join(parts)
175
  backend.mx.synchronize()
176
  generated_at = time.perf_counter()
177
  error = None
178
  answers = None
179
  try:
180
+ answers = parse_generated(text, questions)
181
  if generated.finish_reason != "stop":
182
  error = "Generation reached its token limit without an end-of-answer token"
183
  except ValueError as exc:
184
  error = str(exc)
185
  return {
186
  "answers": answers,
187
+ "raw_text": text,
188
  "valid": error is None,
189
  "error": error,
190
  "finish_reason": generated.finish_reason,
 
211
  return values
212
 
213
 
214
+ def generation_backend(backend):
215
+ """Share evaluated weights while isolating mutable processor state."""
216
+ result = copy(backend)
217
+ if hasattr(backend, "processor"):
218
+ result.processor = deepcopy(backend.processor)
219
+ result.tokenizer = result.processor.tokenizer
220
+ return result
221
+
222
+
223
+ def compare(
224
+ backend,
225
+ state: State,
226
+ questions: dict,
227
+ media_seconds: float,
228
+ emit=None,
229
+ *,
230
+ concurrent=False,
231
+ normal_backend=None,
232
+ ) -> tuple[dict, dict]:
233
+ def evaluate(method, worker_backend, race_started=None):
234
  # No cross-run KV or vision cache. Start each path with completed GPU
235
  # work and a cleared allocator cache; weights remain resident.
236
+ if not concurrent and hasattr(worker_backend, "mx"):
237
+ worker_backend.mx.synchronize()
238
+ worker_backend.mx.clear_cache()
239
+ started = time.perf_counter() if race_started is None else race_started
240
+ if emit:
241
+ emit({"type": "phase_start", "method": method, "media_seconds": media_seconds})
242
+
243
+ def on_answer(path, answer):
244
+ value = discrete_answers({"answer": answer})["answer"]
245
+ if len(path) == 2:
246
+ value = answer["choice"] == "yes"
247
+ emit(
248
+ {
249
+ "type": "answer",
250
+ "method": method,
251
+ "path": list(path),
252
+ "value": value,
253
+ "answer": answer,
254
+ "seconds": media_seconds + time.perf_counter() - started,
255
+ }
256
+ )
257
+
258
+ def on_token(text, tokens):
259
+ emit(
260
+ {
261
+ "type": "token",
262
+ "method": method,
263
+ "text": text,
264
+ "tokens": tokens,
265
+ "seconds": media_seconds + time.perf_counter() - started,
266
+ }
267
+ )
268
+
269
  if method == "parallel":
270
+ engine = DecisionEngine(worker_backend)
271
+ output = (
272
+ engine.system_one(state, questions, on_answer=on_answer)
273
+ if emit
274
+ else engine.system_one(state, questions)
275
+ )
276
+ if hasattr(worker_backend, "mx"):
277
+ worker_backend.mx.synchronize()
278
  output.update(
279
  inference_seconds=time.perf_counter() - started,
280
+ execution=dict(worker_backend.last_stats),
281
  valid=True,
282
  )
283
  else:
284
+ output = (
285
+ generate_answers(worker_backend, state, questions, on_token=on_token)
286
+ if emit
287
+ else generate_answers(worker_backend, state, questions)
288
+ )
289
+ if concurrent:
290
+ output["inference_seconds"] = time.perf_counter() - started
291
  output["total_seconds"] = media_seconds + output["inference_seconds"]
292
+ if emit:
293
+ emit(
294
+ {
295
+ "type": "phase_complete",
296
+ "method": method,
297
+ "seconds": output["total_seconds"],
298
+ "valid": output["valid"],
299
+ }
300
+ )
301
  return output
302
 
303
+ if concurrent:
304
+ # Share evaluated weights only. Tokenizers/processors and KV caches have
305
+ # mutable request state, so normal generation gets its own processor.
306
+ if normal_backend is None:
307
+ normal_backend = generation_backend(backend)
308
+ if hasattr(backend, "mx"):
309
+ backend.mx.synchronize()
310
+ backend.mx.clear_cache()
311
+ ready = Barrier(3)
312
+ race_started = None
313
+
314
+ def worker(method, worker_backend):
315
+ ready.wait()
316
+ mx = getattr(worker_backend, "mx", None)
317
+ stream = mx.new_stream(mx.default_device()) if mx is not None else None
318
+ with mx.stream(stream) if mx is not None else nullcontext():
319
+ try:
320
+ return evaluate(method, worker_backend, race_started)
321
+ finally:
322
+ if mx is not None:
323
+ mx.synchronize(stream)
324
+
325
+ with ThreadPoolExecutor(max_workers=2, thread_name_prefix="comparison") as pool:
326
+ parallel_future = pool.submit(worker, "parallel", backend)
327
+ normal_future = pool.submit(worker, "normal", normal_backend)
328
+ try:
329
+ race_started = time.perf_counter()
330
+ if emit:
331
+ emit({"type": "race_start", "media_seconds": media_seconds})
332
+ ready.wait()
333
+ except BaseException:
334
+ ready.abort()
335
+ raise
336
+ parallel, normal = parallel_future.result(), normal_future.result()
337
+ else:
338
+ parallel = evaluate("parallel", backend)
339
+ normal = evaluate("normal", backend)
340
  decisions = discrete_answers(parallel["answers"])
341
  return parallel, {
342
  "seconds": {"parallel": parallel["total_seconds"], "normal": normal["total_seconds"]},
 
350
  for name in questions
351
  },
352
  "methodology": {
353
+ "execution": "concurrent_shared_gpu" if concurrent else "sequential",
354
  "model": "Same resident Gemma 4 E2B 4-bit weights; float32 compute; all 35 layers",
355
+ "timing": (
356
+ "Both paths start together with one common clock, separate worker streams, processors, and KV caches. They share the same GPU and compete for its resources. This is simultaneous completion time, not isolated throughput."
357
+ if concurrent
358
+ else "One run per path, parallel scorer first, normal generation second. No warm-up runs; first-use effects and run order can affect this observation."
359
+ )
360
+ + " Includes input preparation and inference; shared upload decoding added equally to each path. Excludes model loading, upload transfer, worker setup, and initial allocator reset.",
361
  "cache": "Fresh input KV and media features for every run; ordinary output-token KV caching remains enabled for normal generation.",
362
  "answers": "Normal Gemma generates one JSON object for all fields, greedily, without thinking. Compare choice names, most likely grade levels, and booleans at a 50% threshold. The scorer also returns probability distributions and expected grades. Agreement is not an accuracy measurement.",
363
  },
gemma_rlcd/core.py CHANGED
@@ -216,7 +216,7 @@ class DecisionEngine:
216
  def decide(self, state: State, question: Question) -> dict:
217
  return self.system_one(state, {"answer": question})["answers"]["answer"]
218
 
219
- def system_one(self, state: State, questions: Mapping[str, Question]) -> dict:
220
  if not questions:
221
  raise ValueError("At least one question is required")
222
  jobs = []
@@ -257,8 +257,22 @@ class DecisionEngine:
257
  jobs.append((question_id, child_id, child, criteria))
258
  question_score = getattr(self.backend, "score_questions", None)
259
  batch_score = getattr(self.backend, "score_batch", None)
 
 
 
 
 
 
 
 
 
 
260
  if question_score is not None:
261
- scores = question_score(state, questions)
 
 
 
 
262
  elif batch_score is not None:
263
  scores = batch_score(state, requests)
264
  else:
@@ -267,6 +281,8 @@ class DecisionEngine:
267
  ]
268
  if len(scores) != len(jobs):
269
  raise ValueError("Backend returned the wrong number of question results")
 
 
270
  for (question_id, child_id, question, criteria), score in zip(jobs, scores, strict=True):
271
  probabilities, diagnostics = self._distribution(criteria, score)
272
  if child_id is None:
 
216
  def decide(self, state: State, question: Question) -> dict:
217
  return self.system_one(state, {"answer": question})["answers"]["answer"]
218
 
219
+ def system_one(self, state: State, questions: Mapping[str, Question], on_answer=None) -> dict:
220
  if not questions:
221
  raise ValueError("At least one question is required")
222
  jobs = []
 
257
  jobs.append((question_id, child_id, child, criteria))
258
  question_score = getattr(self.backend, "score_questions", None)
259
  batch_score = getattr(self.backend, "score_batch", None)
260
+
261
+ def completed_scores(rows):
262
+ for index, score in rows:
263
+ question_id, child_id, question, criteria = jobs[index]
264
+ probabilities, diagnostics = self._distribution(criteria, score)
265
+ on_answer(
266
+ (question_id,) if child_id is None else (question_id, child_id),
267
+ self._answer(question, probabilities, diagnostics),
268
+ )
269
+
270
  if question_score is not None:
271
+ scores = (
272
+ question_score(state, questions, on_scores=completed_scores)
273
+ if on_answer
274
+ else question_score(state, questions)
275
+ )
276
  elif batch_score is not None:
277
  scores = batch_score(state, requests)
278
  else:
 
281
  ]
282
  if len(scores) != len(jobs):
283
  raise ValueError("Backend returned the wrong number of question results")
284
+ if on_answer is not None and question_score is None:
285
+ completed_scores(list(enumerate(scores)))
286
  for (question_id, child_id, question, criteria), score in zip(jobs, scores, strict=True):
287
  probabilities, diagnostics = self._distribution(criteria, score)
288
  if child_id is None:
gemma_rlcd/json_backend.py CHANGED
@@ -12,7 +12,7 @@ from .json_scoring import candidate_fields, compile_field
12
  class JSONMLXBackend(CachedMLXBackend):
13
  probability_source = "restricted_json_value_likelihoods"
14
 
15
- def _sequence_scores(self, prefix_cache, prefix_tokens, fields):
16
  """Teacher-force every complete candidate in bounded GPU batches.
17
 
18
  Score all tokens, including string terminators. Shared first tokens and
@@ -26,6 +26,7 @@ class JSONMLXBackend(CachedMLXBackend):
26
  ]
27
  values = {index: [None] * len(field.candidates) for index, field in fields}
28
  batches = []
 
29
  for start in range(0, len(jobs), self.branch_batch_size):
30
  batch = jobs[start : start + self.branch_batch_size]
31
  suffixes = [prefix + candidate[:-1] for _, _, prefix, candidate in batch]
@@ -53,15 +54,20 @@ class JSONMLXBackend(CachedMLXBackend):
53
  for (field_index, choice, _, _), total in zip(batch, totals, strict=True):
54
  values[field_index][choice] = float(total.item())
55
  batches.append(len(batch))
56
- output = {}
57
- for index, field in fields:
58
- scores = tuple(values[index])
59
- mass = min(1.0, sum(math.exp(value) for value in scores))
60
- length = prefix_tokens + len(field.prefix) + max(map(len, field.candidates)) - 1
61
- output[index] = TokenScores(scores, mass, length)
 
 
 
 
 
62
  return output, batches
63
 
64
- def score_questions(self, state, questions):
65
  started = time.perf_counter()
66
  prompt, inputs = prepare_generation(self, state, questions)
67
  fields = [
@@ -102,7 +108,15 @@ class JSONMLXBackend(CachedMLXBackend):
102
  simple = PreparedState(
103
  inputs, [list(field.prefix) for _, field in single], prefix_tokens
104
  )
105
- results = self.branches(simple, prefix_cache, requests)
 
 
 
 
 
 
 
 
106
  batches.extend(self.last_stats["branch_batch_sizes"])
107
  for (index, _), result in zip(single, results, strict=True):
108
  scores[index] = result
@@ -114,7 +128,7 @@ class JSONMLXBackend(CachedMLXBackend):
114
  candidate_batches = []
115
  if multiple:
116
  results, candidate_batches = self._sequence_scores(
117
- prefix_cache, prefix_tokens, multiple
118
  )
119
  for index, result in results.items():
120
  scores[index] = result
 
12
  class JSONMLXBackend(CachedMLXBackend):
13
  probability_source = "restricted_json_value_likelihoods"
14
 
15
+ def _sequence_scores(self, prefix_cache, prefix_tokens, fields, on_scores=None):
16
  """Teacher-force every complete candidate in bounded GPU batches.
17
 
18
  Score all tokens, including string terminators. Shared first tokens and
 
26
  ]
27
  values = {index: [None] * len(field.candidates) for index, field in fields}
28
  batches = []
29
+ output = {}
30
  for start in range(0, len(jobs), self.branch_batch_size):
31
  batch = jobs[start : start + self.branch_batch_size]
32
  suffixes = [prefix + candidate[:-1] for _, _, prefix, candidate in batch]
 
54
  for (field_index, choice, _, _), total in zip(batch, totals, strict=True):
55
  values[field_index][choice] = float(total.item())
56
  batches.append(len(batch))
57
+ completed = []
58
+ for index, field in fields:
59
+ if index in output or any(value is None for value in values[index]):
60
+ continue
61
+ scores = tuple(values[index])
62
+ mass = min(1.0, sum(math.exp(value) for value in scores))
63
+ length = prefix_tokens + len(field.prefix) + max(map(len, field.candidates)) - 1
64
+ output[index] = TokenScores(scores, mass, length)
65
+ completed.append((index, output[index]))
66
+ if on_scores is not None and completed:
67
+ on_scores(completed)
68
  return output, batches
69
 
70
+ def score_questions(self, state, questions, on_scores=None):
71
  started = time.perf_counter()
72
  prompt, inputs = prepare_generation(self, state, questions)
73
  fields = [
 
108
  simple = PreparedState(
109
  inputs, [list(field.prefix) for _, field in single], prefix_tokens
110
  )
111
+
112
+ def completed_batch(start, results):
113
+ on_scores(
114
+ [(single[start + offset][0], score) for offset, score in enumerate(results)]
115
+ )
116
+
117
+ results = self.branches(
118
+ simple, prefix_cache, requests, on_batch=completed_batch if on_scores else None
119
+ )
120
  batches.extend(self.last_stats["branch_batch_sizes"])
121
  for (index, _), result in zip(single, results, strict=True):
122
  scores[index] = result
 
128
  candidate_batches = []
129
  if multiple:
130
  results, candidate_batches = self._sequence_scores(
131
+ prefix_cache, prefix_tokens, multiple, on_scores=on_scores
132
  )
133
  for index, result in results.items():
134
  scores[index] = result
gemma_rlcd/static/demo-utils.js ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ "use strict";
2
+
3
+ const VisualDemo = {
4
+ questions(config, count) {
5
+ if (![32, 64, 128].includes(count)) throw new Error("Choose 32, 64, or 128 checks.");
6
+ return Object.fromEntries(config.groups.map((group) => [group.id, {
7
+ type: "independent",
8
+ instructions: "Which of these checks are visually established?",
9
+ criteria: Object.fromEntries(group.checks.slice(0, count / config.groups.length).map((check) => [check.id, check.description])),
10
+ }]));
11
+ },
12
+ partialBooleans(text) {
13
+ // Read only completed boolean literals in the expected two-level JSON shape.
14
+ // Partial strings, quoted booleans, and malformed suffixes are not answers.
15
+ text = text.replace(/^\s*```(?:json)?\s*\n/, "");
16
+ let index = 0;
17
+ const values = Object.create(null);
18
+ const skip = () => { while (/\s/.test(text[index] || "x")) index++; };
19
+ const take = (char) => { skip(); if (text[index] !== char) return false; index++; return true; };
20
+ const string = () => {
21
+ skip();
22
+ if (text[index] !== '"') return null;
23
+ const start = index++;
24
+ while (index < text.length) {
25
+ if (text[index] === "\\") { index += 2; continue; }
26
+ if (text[index++] === '"') {
27
+ try { return JSON.parse(text.slice(start, index)); } catch { return null; }
28
+ }
29
+ }
30
+ return null;
31
+ };
32
+ if (!take("{")) return values;
33
+ while (index < text.length) {
34
+ const group = string();
35
+ if (group === null || !take(":") || !take("{")) break;
36
+ while (index < text.length) {
37
+ const key = string();
38
+ if (key === null || !take(":")) return values;
39
+ skip();
40
+ const match = /^(true|false)(?=\s*[,}])/.exec(text.slice(index));
41
+ if (!match) return values;
42
+ values[`${group}.${key}`] = match[1] === "true";
43
+ index += match[1].length;
44
+ skip();
45
+ if (take("}")) break;
46
+ if (!take(",")) return values;
47
+ }
48
+ skip();
49
+ if (take("}")) return values;
50
+ if (!take(",")) return values;
51
+ }
52
+ return values;
53
+ },
54
+ };
55
+ if (typeof module !== "undefined") module.exports = VisualDemo;
gemma_rlcd/static/demo.css ADDED
@@ -0,0 +1 @@
 
 
1
+ :root{font-family:-apple-system,BlinkMacSystemFont,"Segoe UI",sans-serif;color:#1d2924;background:#f6f7f4;--green:#23674e;--muted:#69766d;--line:#dfe5dc}*{box-sizing:border-box}body{margin:0;font-size:14px;line-height:1.5}header{display:flex;align-items:center;justify-content:space-between;gap:20px;padding:20px 36px;border-bottom:1px solid var(--line);background:#fcfdfb}.brand{font-weight:650;font-size:17px;color:inherit;text-decoration:none}nav{display:flex;gap:24px;align-items:center;font-size:12px}a{color:var(--green)}#model-status{color:var(--muted)}main{max-width:1560px;margin:auto;padding:35px 36px 70px}h1,h2,h3,p{margin:0}h1{font-size:36px;letter-spacing:-1.4px;font-weight:580;line-height:1.25;margin:8px 0}h1>span{color:var(--green)}h2{font-size:15px;font-weight:650}h3{font-size:14px;font-weight:620}.eyebrow{font-size:10px;font-weight:650;letter-spacing:1.7px;color:var(--muted)}.heading{display:flex;align-items:center;justify-content:space-between;gap:20px;margin-bottom:28px}.heading p{color:var(--muted);font-size:13px}.workspace{display:grid;grid-template-columns:minmax(280px,.75fr) minmax(0,1.6fr);gap:24px;align-items:start}.input-panel,.live-panel,.wall{border:1px solid var(--line);background:white;border-radius:10px;padding:22px}.section-title{display:flex;align-items:center;justify-content:space-between;gap:12px;margin-bottom:16px}.small{font-size:11px;color:var(--muted);line-height:1.65}.preview{height:260px;background:#eff2ed;border-radius:6px;overflow:hidden;display:grid;place-items:center;cursor:pointer}.preview img,.preview video{width:100%;height:100%;max-height:260px;object-fit:contain}.preview.dragging{outline:3px solid #7aab8f}.media-controls{display:flex;gap:12px;align-items:center;margin-top:12px}.file-button{font-size:12px;font-weight:550;color:var(--green);cursor:pointer;position:relative;flex-shrink:0}.file-button input{position:absolute;inset:0;opacity:0;width:100%;cursor:pointer}#file-name{font-size:10px;color:var(--muted);overflow:hidden;text-overflow:ellipsis;white-space:nowrap}#credit{margin-top:5px;min-height:18px}label:not(.file-button){display:block;font-size:11px;font-weight:550;margin:15px 0 6px}button,input,textarea,select{font:inherit;color:inherit}button,select{cursor:pointer}button{border:1px solid transparent;border-radius:6px;padding:10px 14px;font-size:12px;font-weight:550}button:disabled{opacity:.5;cursor:default}.primary{background:var(--green);color:white}.primary:hover:not(:disabled){background:#18513b}.secondary{background:white;border-color:var(--line)}.text-button{padding:2px;background:none;color:var(--green);font-size:11px}.input-panel select,.input-panel textarea{width:100%;border:1px solid var(--line);border-radius:6px;padding:10px;background:#fcfdfb;font-size:12px}.input-panel textarea{resize:vertical;line-height:1.5;margin-bottom:9px}.actions{display:flex;gap:10px;margin:18px 0 9px}.actions .primary{flex:1}.engines{display:grid;grid-template-columns:1fr 1fr;gap:24px}.engine{min-width:0}.engine-title{display:flex;align-items:center;justify-content:space-between;gap:8px}.engine-title>span{font-size:10px;color:var(--muted)}.parallel h3,.parallel .clock{color:var(--green)}.clock{display:block;font-size:50px;font-weight:450;letter-spacing:-2px;font-variant-numeric:tabular-nums;margin-top:8px}.clock>span{font-size:15px;letter-spacing:0;color:var(--muted);margin-left:7px}.progress{height:4px;background:#edf0e9;border-radius:3px;overflow:hidden;margin-top:12px}.progress i{display:block;height:100%;width:0;background:var(--green)}.normal .progress i{background:#4b5850}.engine-stats{display:flex;justify-content:space-between;gap:5px;color:var(--muted);font-size:10px;margin-top:7px}.engine>.small{margin-top:10px}.verdict{border-top:1px solid var(--line);border-bottom:1px solid var(--line);margin:22px 0;padding:15px 0;min-height:54px;font-size:13px;color:var(--green)}.verdict strong{font-size:22px;font-weight:600;margin-right:8px}.streams{display:grid;grid-template-columns:1fr 1fr;gap:20px;min-width:0}.streams>div{min-width:0}.stream-label{display:flex;justify-content:space-between;gap:10px;font-size:11px;color:var(--muted)}.live-dot{font-size:8px;letter-spacing:1px;color:var(--green)}pre{height:235px;margin:10px 0 0;background:#f6f8f3;border:1px solid #e9eee4;border-radius:6px;padding:12px;white-space:pre-wrap;overflow-wrap:anywhere;overflow:auto;font:11px/1.7 ui-monospace,SFMono-Regular,Menlo,monospace;color:#3e6550}.normal-stream{color:#46504a}.method-note{font-size:10px;color:var(--muted);line-height:1.8;margin-top:17px}.wall{margin-top:24px}.wall-heading{display:flex;justify-content:space-between;align-items:center;gap:20px;margin-bottom:20px}.wall-heading p{margin-top:5px}.filters{display:flex;align-items:center;gap:12px;flex-shrink:0}.filters select{font-size:11px;padding:8px;border:1px solid var(--line);border-radius:6px;background:white}.group{margin-top:19px}.group:first-child{margin-top:0}.group-title{font-size:11px;color:var(--muted);font-weight:550;margin-bottom:9px}.check-grid{display:grid;grid-template-columns:repeat(8,minmax(0,1fr));gap:7px}.check{border:1px solid #e3e8de;border-radius:5px;padding:8px 9px;min-width:0;min-height:59px;transition:background .2s,border-color .2s}.check-name{display:block;font-size:10px;white-space:nowrap;overflow:hidden;text-overflow:ellipsis;color:#53624f}.values{display:flex;gap:12px;justify-content:space-between;margin-top:6px;font:10px ui-monospace,monospace;color:#8a9585}.values b{font-weight:500}.check.detected{background:#f0f7ec;border-color:#c7dcc1}.check.different{background:#fff6e9;border-color:#e9c998}.values .yes{color:#23674e}.values .no{color:#727d70}.check.flash{animation:arrive .55s ease-out}.details{font-size:11px;color:var(--muted);margin-top:20px;max-width:1000px}.details p{margin-top:10px}.details summary{cursor:pointer}#error{margin-top:12px;padding:11px;border:1px solid #eccfc4;background:#fff6f0;color:#963e2e;font-size:12px;border-radius:6px}.sr-only{position:absolute;width:1px;height:1px;overflow:hidden;clip:rect(0,0,0,0)}[hidden]{display:none!important}:focus-visible{outline:3px solid #75a58c;outline-offset:3px}@keyframes arrive{0%{background:#dceccd}100%{}}@media(prefers-reduced-motion:reduce){*{animation:none!important;transition:none!important}}@media(max-width:1200px){.check-grid{grid-template-columns:repeat(6,minmax(0,1fr))}.workspace{grid-template-columns:minmax(275px,.8fr) minmax(0,1.4fr)}.clock{font-size:43px}.engine-title{align-items:flex-start;flex-direction:column;gap:2px}}@media(max-width:850px){header{padding:17px 20px}main{padding:25px 20px 50px}.workspace{grid-template-columns:1fr}.preview{height:280px}.preview img,.preview video{max-height:280px}.input-panel{display:grid;grid-template-columns:1fr 1fr;column-gap:20px}.input-panel>*{grid-column:1/-1}.check-grid{grid-template-columns:repeat(4,minmax(0,1fr))}.engine-title{flex-direction:row}.wall-heading{align-items:flex-start;flex-direction:column}h1{font-size:32px}}@media(max-width:520px){header{align-items:flex-start}.brand{font-size:14px}nav{flex-direction:column;gap:3px;align-items:flex-end;font-size:10px}main{padding:22px 13px 40px}.heading{align-items:flex-start}.heading>button{padding:8px;font-size:10px;white-space:nowrap}h1{font-size:27px}.heading p{font-size:12px}.input-panel,.live-panel,.wall{padding:16px}.engines{gap:17px}.engine-title{align-items:flex-start;flex-direction:column}.clock{font-size:40px}.streams{gap:12px}pre{height:215px;font-size:10px;padding:9px}.check-grid{grid-template-columns:repeat(3,minmax(0,1fr))}.check{padding:7px}.values{gap:6px}.preview{height:230px}.preview img,.preview video{max-height:230px}.filters{flex-wrap:wrap}.wall-heading p{font-size:10px}}
gemma_rlcd/static/demo.html ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <!doctype html>
2
+ <html lang="en">
3
+ <head>
4
+ <meta charset="utf-8"><meta name="viewport" content="width=device-width,initial-scale=1">
5
+ <title>Live visual demo · Gemma E2B RLCD</title>
6
+ <link rel="stylesheet" href="/static/demo.css"><script src="/static/demo-utils.js" defer></script><script src="/static/demo.js" defer></script>
7
+ </head>
8
+ <body>
9
+ <header><a class="brand" href="/demo">Gemma E2B RLCD</a><nav><span id="model-status" role="status">Connecting…</span><a href="/">Open playground ↗</a></nav></header>
10
+ <main>
11
+ <div class="heading"><div><span class="eyebrow">LIVE MULTIMODAL COMPARISON</span><h1>One scene. <span id="headline-count">128</span> decisions.</h1><p>Both start together. Watch batched decisions race against normal Gemma’s streamed JSON.</p></div><button id="export" class="secondary" disabled>Export run ↓</button></div>
12
+ <div class="workspace">
13
+ <section class="input-panel" aria-labelledby="input-title">
14
+ <div class="section-title"><h2 id="input-title">The evidence</h2><button id="sample" class="text-button">Use sample photo</button></div>
15
+ <div id="dropzone" class="preview" tabindex="0" role="button" aria-label="Choose an image or video"><img id="image-preview" alt="Selected input"><video id="video-preview" controls playsinline hidden></video><div id="empty-preview" hidden>Drop an image or video</div></div>
16
+ <div class="media-controls"><label class="file-button">Choose image or video<input id="media-file" type="file" accept="image/png,image/jpeg,image/webp,image/bmp,video/*"></label><span id="file-name"></span></div>
17
+ <p id="credit" class="small"></p>
18
+ <label for="output-count">Number of visual checks</label><select id="output-count"><option value="32">32 checks · quick comparison</option><option value="64">64 checks · larger output</option><option value="128" selected>128 checks · full visual inventory</option></select>
19
+ <label for="instructions">Instructions</label><textarea id="instructions" rows="3" placeholder="Optional instructions for both models"></textarea>
20
+ <p class="small">One image or video. Silent clips up to 60 s; clips with audio up to 30 s. Sampled video frames and soundtrack go to both paths.</p>
21
+ <div class="actions"><button id="run" class="primary" disabled>Compare · stream answers</button><button id="stop" class="secondary" hidden>Stop</button></div>
22
+ <p id="run-note" class="small" role="status">The model loads once. Every comparison uses fresh input state.</p>
23
+ <div id="error" role="alert" hidden></div>
24
+ </section>
25
+ <section class="live-panel" aria-labelledby="live-title">
26
+ <div class="section-title"><h2 id="live-title">Live output</h2><span id="run-phase" class="small">Ready</span></div>
27
+ <div class="engines">
28
+ <article class="engine parallel"><div class="engine-title"><h3>Parallel scorer</h3><span id="parallel-state">Waiting</span></div><strong class="clock" id="parallel-clock">0.00<span>s</span></strong><div class="progress"><i id="parallel-progress"></i></div><div class="engine-stats"><span id="parallel-count">0 / 128 decisions</span><span>GPU batches</span></div><p id="parallel-first" class="small">First answer —</p></article>
29
+ <article class="engine normal"><div class="engine-title"><h3>Normal Gemma</h3><span id="normal-state">Waiting</span></div><strong class="clock" id="normal-clock">0.00<span>s</span></strong><div class="progress"><i id="normal-progress"></i></div><div class="engine-stats"><span id="normal-count">0 / 128 decisions</span><span id="token-count">0 tokens</span></div><p id="normal-first" class="small">First answer —</p></article>
30
+ </div>
31
+ <div id="verdict" class="verdict" aria-live="polite">Upload a scene or try the sample. Results are measured live.</div>
32
+ <div class="streams">
33
+ <div><div class="stream-label"><span>Completed decisions</span><span class="live-dot">LIVE</span></div><pre id="parallel-stream" aria-label="Parallel scorer output">Answers appear as soon as each batch completes.</pre></div>
34
+ <div><div class="stream-label"><span>Generated JSON</span><span class="live-dot">LIVE</span></div><pre id="normal-stream" aria-label="Normal Gemma generated JSON">Real token output will stream here.</pre></div>
35
+ </div>
36
+ <p class="method-note">Both paths start together on the same GPU, with independent input state and one shared clock. Same weights, media, questions, and compact answer contract. Timings include resource contention; model loading excluded.</p>
37
+ </section>
38
+ </div>
39
+ <section class="wall" aria-labelledby="wall-title"><div class="wall-heading"><div><h2 id="wall-title">Every decision, as it lands</h2><p class="small">A check is “yes” when visually established. P = parallel probability of yes · G = Gemma’s boolean. Streaming JSON is provisional until validated.</p></div><div class="filters"><label for="filter" class="sr-only">Filter decisions</label><select id="filter"><option value="all">All checks</option><option value="yes">Detected by either</option><option value="different">Different answers</option></select><span id="agreement" class="small"></span></div></div><div id="output-grid"></div><p id="no-matches" class="small" hidden>No completed checks match this filter.</p></section>
40
+ <details class="details"><summary>What is being measured?</summary><p>Both paths process the same complete image or sampled video, instructions, and label descriptions. The parallel scorer returns probabilities as field batches finish. Normal Gemma is asked to generate one compact JSON object with boolean decisions, without explanations or probability prose.</p><p>The timers include input preparation and inference, with shared upload decoding added equally. They exclude model loading, upload transfer, and allocator reset. This is one simultaneous run per path; GPU contention and first-use effects can affect timing. The result measures completion time while both are running, not isolated throughput. Matching answers measures agreement, not correctness. Invalid normal output is shown without a speedup claim.</p><p>Video processing targets one frame per second with a 32-frame cap. Long media plus 128 questions can exceed the 8,192-token input limit; reduce the number of checks or use a shorter clip. Media is never silently trimmed.</p></details>
41
+ </main>
42
+ </body>
43
+ </html>
gemma_rlcd/static/demo.js ADDED
@@ -0,0 +1,242 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ "use strict";
2
+ const $ = (selector) => document.querySelector(selector);
3
+ const escapeHTML = (value) => String(value).replace(/[&<>"']/g, (char) => ({"&":"&amp;","<":"&lt;",">":"&gt;",'"':"&quot;","'":"&#39;"}[char]));
4
+ let config, media, mediaURL, active = false, ready = false, controller, finalResult = null;
5
+ let checks = [], rows = new Map(), completed = {parallel: new Map(), normal: new Map()};
6
+ let clocks = {}, rawText = "", parallelLines = [], events = [], firstAnswer = {}, frame = null, sampleVersion = 0;
7
+
8
+ function error(message) { $("#error").textContent = message; $("#error").hidden = !message; }
9
+ function seconds(value) { return `${value.toFixed(2)}<span>s</span>`; }
10
+ function count() { return Number($("#output-count").value); }
11
+ function setMedia(file, credit = "") {
12
+ if (active) return;
13
+ if (file.size > 200 * 1024 * 1024) return error("The media file exceeds 200 MB.");
14
+ const extension = file.name.split(".").pop().toLowerCase();
15
+ const kind = file.type.startsWith("image/") || ["jpg","jpeg","png","webp","bmp"].includes(extension) ? "image" : file.type.startsWith("video/") || ["mp4","mov","webm","mkv","m4v"].includes(extension) ? "video" : null;
16
+ if (!kind) return error("Choose an image or video file.");
17
+ if (mediaURL) URL.revokeObjectURL(mediaURL);
18
+ media = {file, kind}; mediaURL = URL.createObjectURL(file);
19
+ $("#image-preview").hidden = kind !== "image";
20
+ $("#video-preview").hidden = kind !== "video";
21
+ $("#image-preview").removeAttribute("src");
22
+ $("#video-preview").removeAttribute("src");
23
+ $(`#${kind}-preview`).src = mediaURL;
24
+ $("#image-preview").alt = file.name;
25
+ $("#file-name").textContent = file.name;
26
+ $("#credit").innerHTML = credit;
27
+ reset(); error("");
28
+ }
29
+ async function sample() {
30
+ const version = ++sampleVersion;
31
+ $("#sample").disabled = true;
32
+ try {
33
+ const response = await fetch(config.sample.url);
34
+ if (!response.ok) throw new Error("Could not load the sample photo.");
35
+ const blob = await response.blob();
36
+ if (version !== sampleVersion || active) return;
37
+ setMedia(new File([blob], "times-square.jpg", {type:"image/jpeg"}), `<a href="${escapeHTML(config.sample.source)}" target="_blank" rel="noreferrer">${escapeHTML(config.sample.author)} · ${escapeHTML(config.sample.license)}</a>`);
38
+ } catch (failure) { error(failure.message); }
39
+ finally { $("#sample").disabled = active; }
40
+ }
41
+ function reset() {
42
+ finalResult = null; rawText = ""; parallelLines = []; events = []; firstAnswer = {}; clocks = {};
43
+ completed = {parallel: new Map(), normal: new Map()};
44
+ $("#headline-count").textContent = count();
45
+ $("#parallel-stream").textContent = "Answers appear as soon as each batch completes.";
46
+ $("#normal-stream").textContent = "Real token output will stream here.";
47
+ $("#token-count").textContent = "0 tokens";
48
+ $("#run-phase").textContent = "Ready";
49
+ $("#agreement").textContent = "";
50
+ $("#verdict").textContent = "Upload a scene or try the sample. Results are measured live.";
51
+ $("#export").disabled = true;
52
+ for (const method of ["parallel","normal"]) {
53
+ $(`#${method}-clock`).innerHTML = seconds(0);
54
+ $(`#${method}-state`).textContent = "Waiting";
55
+ $(`#${method}-first`).textContent = "First answer —";
56
+ $(`#${method}-progress`).style.width = "0%";
57
+ $(`#${method}-count`).textContent = `0 / ${count()} decisions`;
58
+ }
59
+ renderGrid(); syncButtons();
60
+ }
61
+ function renderGrid() {
62
+ checks = [];
63
+ $("#output-grid").innerHTML = config.groups.map((group) => {
64
+ const items = group.checks.slice(0, count() / config.groups.length);
65
+ return `<section class="group"><h3 class="group-title">${escapeHTML(group.name)} · ${items.length} checks</h3><div class="check-grid">${items.map((check) => {
66
+ const path = `${group.id}.${check.id}`; checks.push(path);
67
+ return `<div class="check" data-path="${escapeHTML(path)}" title="${escapeHTML(check.description)}"><span class="check-name">${escapeHTML(check.label)}</span><div class="values"><span>P <b data-method="parallel">—</b></span><span>G <b data-method="normal">—</b></span></div></div>`;
68
+ }).join("")}</div></section>`;
69
+ }).join("");
70
+ rows = new Map(Array.from(document.querySelectorAll(".check"), (row) => [row.dataset.path, row]));
71
+ filter();
72
+ }
73
+ function updateDecision(method, path, value, probability, elapsed) {
74
+ const row = rows.get(path);
75
+ if (!row || completed[method].has(path) && completed[method].get(path).value === value) return;
76
+ completed[method].set(path, {value, probability});
77
+ const cell = row.querySelector(`[data-method="${method}"]`);
78
+ cell.textContent = method === "parallel" ? `${Math.round(probability * 100)}%` : value ? "yes" : "no";
79
+ cell.className = value ? "yes" : "no";
80
+ row.classList.add("flash"); setTimeout(() => row.classList.remove("flash"), 600);
81
+ const parallel = completed.parallel.get(path), normal = completed.normal.get(path);
82
+ row.classList.toggle("detected", Boolean(parallel?.value || normal?.value));
83
+ row.classList.toggle("different", Boolean(parallel && normal && parallel.value !== normal.value));
84
+ if (firstAnswer[method] === undefined) {
85
+ firstAnswer[method] = elapsed;
86
+ $(`#${method}-first`).textContent = `First answer ${elapsed.toFixed(2)} s`;
87
+ }
88
+ $(`#${method}-count`).textContent = `${completed[method].size} / ${count()} decisions`;
89
+ $(`#${method}-progress`).style.width = `${100 * completed[method].size / count()}%`;
90
+ filter();
91
+ }
92
+ function filter() {
93
+ const mode = $("#filter").value;
94
+ let visible = 0;
95
+ for (const [path, row] of rows) {
96
+ const p = completed.parallel.get(path), n = completed.normal.get(path);
97
+ row.hidden = mode === "yes" ? !(p?.value || n?.value) : mode === "different" ? !(p && n && p.value !== n.value) : false;
98
+ if (!row.hidden) visible++;
99
+ }
100
+ for (const group of document.querySelectorAll(".group")) group.hidden = !Array.from(group.querySelectorAll(".check")).some((row) => !row.hidden);
101
+ $("#no-matches").hidden = visible > 0;
102
+ }
103
+ function tick() {
104
+ for (const [method, clock] of Object.entries(clocks)) {
105
+ if (clock.running) $(`#${method}-clock`).innerHTML = seconds((performance.now() - clock.start) / 1000);
106
+ }
107
+ if (active) frame = requestAnimationFrame(tick);
108
+ }
109
+ function handle(event) {
110
+ events.push(event.type === "complete" ? {type: "complete"} : event);
111
+ const method = event.method;
112
+ if (event.type === "accepted") { $("#run-phase").textContent = "Preparing media"; return; }
113
+ if (event.type === "error") throw new Error(event.error);
114
+ if (event.type === "race_start") {
115
+ const start = performance.now() - event.media_seconds * 1000;
116
+ for (const name of ["parallel","normal"]) {
117
+ clocks[name] = {running:true, start};
118
+ $(`#${name}-state`).textContent = "Running";
119
+ }
120
+ $("#run-phase").textContent = "Both running · live";
121
+ $("#verdict").textContent = "Both paths started together. Watch the answers arrive…";
122
+ } else if (event.type === "phase_start") {
123
+ if (!clocks[method]) clocks[method] = {running:true, start:performance.now() - event.media_seconds * 1000};
124
+ $(`#${method}-state`).textContent = "Running";
125
+ } else if (event.type === "answer") {
126
+ const path = event.path.join(".");
127
+ const probability = event.answer.probabilities.yes;
128
+ updateDecision("parallel",path,event.value,probability,event.seconds);
129
+ parallelLines.push(`${path}: ${event.value} (${(probability*100).toFixed(1)}% yes)`);
130
+ $("#parallel-stream").textContent = parallelLines.join("\n");
131
+ $("#parallel-stream").scrollTop = $("#parallel-stream").scrollHeight;
132
+ } else if (event.type === "token") {
133
+ rawText += event.text;
134
+ $("#normal-stream").textContent = rawText;
135
+ $("#normal-stream").scrollTop = $("#normal-stream").scrollHeight;
136
+ $("#token-count").textContent = `${event.tokens} tokens`;
137
+ for (const [path,value] of Object.entries(VisualDemo.partialBooleans(rawText))) updateDecision("normal",path,value,null,event.seconds);
138
+ } else if (event.type === "phase_complete") {
139
+ clocks[method].running = false;
140
+ $(`#${method}-clock`).innerHTML = seconds(event.seconds);
141
+ $(`#${method}-state`).textContent = event.valid ? "Complete" : "Invalid output";
142
+ const other = method === "parallel" ? "normal" : "parallel";
143
+ if (clocks[other]?.running) {
144
+ const name = method === "parallel" ? "Parallel scorer" : "Normal Gemma";
145
+ const otherName = other === "parallel" ? "Parallel scorer" : "Normal Gemma";
146
+ $("#run-phase").textContent = `${otherName} still running`;
147
+ $("#verdict").textContent = `${name} ${event.valid ? "finished" : "returned invalid output"} in ${event.seconds.toFixed(2)} s. ${otherName} is still working…`;
148
+ }
149
+ } else if (event.type === "complete") {
150
+ finalResult = event.result;
151
+ const comparison = finalResult.comparison;
152
+ let matched = 0;
153
+ if (comparison.normal.valid) {
154
+ for (const [group, values] of Object.entries(comparison.normal.answers)) {
155
+ for (const [key,value] of Object.entries(values)) {
156
+ const path = `${group}.${key}`;
157
+ updateDecision("normal",path,value,null,comparison.seconds.normal);
158
+ if (completed.parallel.get(path)?.value === value) matched++;
159
+ }
160
+ }
161
+ $("#agreement").textContent = `${matched} / ${count()} matched`;
162
+ const ratio = comparison.normal_over_parallel;
163
+ $("#verdict").innerHTML = `<strong>${ratio >= 1 ? ratio.toFixed(2) : (1 / ratio).toFixed(2)}× ${ratio >= 1 ? "faster" : "slower"}</strong> this run · ${matched} / ${count()} matching answers`;
164
+ } else {
165
+ $("#verdict").textContent = "Normal Gemma returned an invalid answer. No valid-response speedup is reported.";
166
+ error(comparison.normal.error);
167
+ $("#agreement").textContent = "Normal output invalid";
168
+ }
169
+ $("#run-phase").textContent = "Comparison complete";
170
+ $("#export").disabled = false;
171
+ }
172
+ }
173
+ function syncButtons() {
174
+ $("#run").disabled = !ready || active || !media;
175
+ $("#stop").hidden = !active;
176
+ for (const selector of ["#sample","#media-file","#output-count","#instructions"]) $(selector).disabled = active;
177
+ $("#run").textContent = active ? "Streaming…" : "Compare · stream answers";
178
+ }
179
+ async function run() {
180
+ if (active || !ready || !media) return;
181
+ reset(); error(""); active = true; controller = new AbortController(); syncButtons(); tick();
182
+ const spec = {text:"", instructions:$("#instructions").value, questions:VisualDemo.questions(config,count()), media:[{name:media.file.name,kind:media.kind}]};
183
+ const form = new FormData(); form.append("spec",JSON.stringify(spec)); form.append("media",media.file);
184
+ try {
185
+ const response = await fetch("/api/compare-stream",{method:"POST",body:form,signal:controller.signal});
186
+ if (!response.ok) { const data=await response.json(); throw new Error(data.error || "Comparison failed."); }
187
+ const reader=response.body.getReader(), decoder=new TextDecoder(); let pending="";
188
+ while (true) {
189
+ const {value,done}=await reader.read(); pending += decoder.decode(value,{stream:!done});
190
+ let end;
191
+ while ((end=pending.indexOf("\n"))>=0) { const line=pending.slice(0,end); pending=pending.slice(end+1); if(line.trim()) handle(JSON.parse(line)); }
192
+ if(done) break;
193
+ }
194
+ if(pending.trim()) handle(JSON.parse(pending));
195
+ if(!finalResult) throw new Error("The stream ended before the comparison completed.");
196
+ finalResult = {request:spec,response:finalResult,stream_events:events,first_answer_seconds:firstAnswer};
197
+ $("#run-note").textContent = "Complete. Export preserves the answers, events, and measured timings.";
198
+ } catch(failure) {
199
+ controller.abort();
200
+ const cancelled=failure.name==="AbortError";
201
+ if(!cancelled) error(failure.message);
202
+ $("#run-phase").textContent=cancelled ? "Stopped" : "Run failed";
203
+ $("#verdict").textContent=cancelled ? "Stopped. Partial results are shown; no complete-response comparison." : "The comparison did not complete.";
204
+ for(const [method,clock] of Object.entries(clocks)) if(clock.running) $(`#${method}-state`).textContent=cancelled ? "Stopped" : "Failed";
205
+ } finally {
206
+ active=false; cancelAnimationFrame(frame); for(const clock of Object.values(clocks)) clock.running=false;
207
+ syncButtons(); status();
208
+ }
209
+ }
210
+ async function status() {
211
+ try {
212
+ const response=await fetch("/api/status"), state=await response.json();
213
+ ready=state.ready && !state.busy;
214
+ $("#model-status").textContent=state.error ? "Model unavailable" : state.busy ? "Gemma · running" : state.ready ? "Gemma 4 E2B · ready" : "Loading model…";
215
+ if(state.error) error(state.error);
216
+ } catch { ready=false; $("#model-status").textContent="Server disconnected"; }
217
+ syncButtons();
218
+ }
219
+ $("#run").addEventListener("click",run);
220
+ $("#stop").addEventListener("click",()=>controller?.abort());
221
+ $("#sample").addEventListener("click",sample);
222
+ $("#filter").addEventListener("change",filter);
223
+ $("#output-count").addEventListener("change",reset);
224
+ $("#instructions").addEventListener("input",()=>{ if(finalResult) $("#run-note").textContent="Instructions changed. Run again to update the results."; });
225
+ $("#media-file").addEventListener("change",(event)=>{ sampleVersion++; if(event.target.files[0]) setMedia(event.target.files[0]); event.target.value=""; });
226
+ $("#dropzone").addEventListener("click",(event)=>{ if(event.target.tagName!=="VIDEO" && !active) $("#media-file").click(); });
227
+ $("#dropzone").addEventListener("keydown",(event)=>{ if(["Enter"," "].includes(event.key) && !active) {event.preventDefault(); $("#media-file").click();} });
228
+ for(const name of ["dragover","dragenter"]) $("#dropzone").addEventListener(name,(event)=>{event.preventDefault(); if(!active) $("#dropzone").classList.add("dragging");});
229
+ for(const name of ["dragleave","drop"]) $("#dropzone").addEventListener(name,(event)=>{event.preventDefault(); $("#dropzone").classList.remove("dragging");});
230
+ $("#dropzone").addEventListener("drop",(event)=>{sampleVersion++; if(event.dataTransfer.files[0]) setMedia(event.dataTransfer.files[0]);});
231
+ $("#export").addEventListener("click",()=>{
232
+ if(!finalResult) return;
233
+ const url=URL.createObjectURL(new Blob([JSON.stringify(finalResult,null,2)],{type:"application/json"}));
234
+ const link=document.createElement("a"); link.href=url; link.download="gemma-visual-comparison.json"; link.click(); setTimeout(()=>URL.revokeObjectURL(url),1000);
235
+ });
236
+ async function init() {
237
+ try {
238
+ const response=await fetch("/static/visual-demo.json"); if(!response.ok) throw new Error("Could not load the visual checks.");
239
+ config=await response.json(); $("#instructions").value=config.instructions; reset(); await sample(); await status(); setInterval(status,2000);
240
+ } catch(failure) { error(failure.message); }
241
+ }
242
+ init();
gemma_rlcd/static/index.html CHANGED
@@ -12,7 +12,7 @@
12
  <body>
13
  <header class="topbar">
14
  <a class="brand" href="/" aria-label="Gemma E2B RLCD home"><span class="brand-mark" aria-hidden="true"><i></i><i></i><i></i></span>Gemma E2B RLCD<span class="local-tag">LOCAL</span></a>
15
- <div class="model-status"><span id="status-dot" class="status-dot loading"></span><span id="model-status" role="status">Loading Gemma 4 E2B…</span></div>
16
  </header>
17
  <main>
18
  <div class="page-heading">
 
12
  <body>
13
  <header class="topbar">
14
  <a class="brand" href="/" aria-label="Gemma E2B RLCD home"><span class="brand-mark" aria-hidden="true"><i></i><i></i><i></i></span>Gemma E2B RLCD<span class="local-tag">LOCAL</span></a>
15
+ <div class="model-status"><a href="/demo">Live visual demo ↗</a><span id="status-dot" class="status-dot loading"></span><span id="model-status" role="status">Loading Gemma 4 E2B…</span></div>
16
  </header>
17
  <main>
18
  <div class="page-heading">
gemma_rlcd/static/sample-street.jpg ADDED

Git LFS Details

  • SHA256: 61eb9cee360a7e5efa4c6e70164d9e5ee484a20c012f014f58a4bf4f249482f2
  • Pointer size: 132 Bytes
  • Size of remote file: 5.41 MB
gemma_rlcd/static/visual-demo.json ADDED
@@ -0,0 +1,677 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "title": "Visual inventory",
3
+ "instructions": "Evaluate only visible evidence. For video, a check is true if visible in at least one sampled frame. A spoken mention or a printed picture of an object does not establish that the real object is present. If a check is not visually established, answer false.",
4
+ "groups": [
5
+ {
6
+ "id": "people",
7
+ "name": "People & activity",
8
+ "checks": [
9
+ {
10
+ "id": "person",
11
+ "label": "Person",
12
+ "description": "A person is visible"
13
+ },
14
+ {
15
+ "id": "crowd",
16
+ "label": "Crowd",
17
+ "description": "At least ten people are visible"
18
+ },
19
+ {
20
+ "id": "walking",
21
+ "label": "Walking",
22
+ "description": "A person is visibly walking"
23
+ },
24
+ {
25
+ "id": "sitting",
26
+ "label": "Sitting",
27
+ "description": "A person is sitting"
28
+ },
29
+ {
30
+ "id": "standing",
31
+ "label": "Standing",
32
+ "description": "A person is standing"
33
+ },
34
+ {
35
+ "id": "running",
36
+ "label": "Running",
37
+ "description": "A person is running"
38
+ },
39
+ {
40
+ "id": "cycling",
41
+ "label": "Cycling",
42
+ "description": "A person is riding a bicycle"
43
+ },
44
+ {
45
+ "id": "carrying_bag",
46
+ "label": "Carrying bag",
47
+ "description": "A person is carrying a bag"
48
+ },
49
+ {
50
+ "id": "backpack",
51
+ "label": "Backpack",
52
+ "description": "A person is wearing a backpack"
53
+ },
54
+ {
55
+ "id": "handbag",
56
+ "label": "Handbag",
57
+ "description": "A handbag is visible"
58
+ },
59
+ {
60
+ "id": "hat",
61
+ "label": "Hat",
62
+ "description": "A person is wearing a hat or cap"
63
+ },
64
+ {
65
+ "id": "sunglasses",
66
+ "label": "Sunglasses",
67
+ "description": "A person is wearing sunglasses"
68
+ },
69
+ {
70
+ "id": "umbrella",
71
+ "label": "Umbrella",
72
+ "description": "A person is holding an umbrella"
73
+ },
74
+ {
75
+ "id": "phone_in_hand",
76
+ "label": "Phone in hand",
77
+ "description": "A person is holding a phone"
78
+ },
79
+ {
80
+ "id": "using_camera",
81
+ "label": "Using camera",
82
+ "description": "A person is using a camera"
83
+ },
84
+ {
85
+ "id": "raised_hand",
86
+ "label": "Raised hand",
87
+ "description": "A person has a raised hand"
88
+ },
89
+ {
90
+ "id": "waving",
91
+ "label": "Waving",
92
+ "description": "A person is waving"
93
+ },
94
+ {
95
+ "id": "pointing",
96
+ "label": "Pointing",
97
+ "description": "A person is pointing"
98
+ },
99
+ {
100
+ "id": "eating",
101
+ "label": "Eating",
102
+ "description": "A person is eating"
103
+ },
104
+ {
105
+ "id": "drinking",
106
+ "label": "Drinking",
107
+ "description": "A person is drinking"
108
+ },
109
+ {
110
+ "id": "stroller",
111
+ "label": "Stroller",
112
+ "description": "A baby stroller is visible"
113
+ },
114
+ {
115
+ "id": "walking_dog",
116
+ "label": "Walking dog",
117
+ "description": "A person is walking a dog on a leash"
118
+ },
119
+ {
120
+ "id": "uniform",
121
+ "label": "Uniform",
122
+ "description": "A person is wearing a uniform"
123
+ },
124
+ {
125
+ "id": "helmet",
126
+ "label": "Helmet",
127
+ "description": "A person is wearing a helmet"
128
+ },
129
+ {
130
+ "id": "high_vis_vest",
131
+ "label": "High vis vest",
132
+ "description": "A high-visibility safety vest is visible"
133
+ },
134
+ {
135
+ "id": "face_mask",
136
+ "label": "Face mask",
137
+ "description": "A person is wearing a face mask"
138
+ },
139
+ {
140
+ "id": "red_clothing",
141
+ "label": "Red clothing",
142
+ "description": "Red clothing is visible"
143
+ },
144
+ {
145
+ "id": "blue_clothing",
146
+ "label": "Blue clothing",
147
+ "description": "Blue clothing is visible"
148
+ },
149
+ {
150
+ "id": "white_clothing",
151
+ "label": "White clothing",
152
+ "description": "White clothing is visible"
153
+ },
154
+ {
155
+ "id": "striped_clothing",
156
+ "label": "Striped clothing",
157
+ "description": "Striped clothing is visible"
158
+ },
159
+ {
160
+ "id": "group_interaction",
161
+ "label": "Group interaction",
162
+ "description": "Two people are visibly interacting"
163
+ },
164
+ {
165
+ "id": "person_lying_down",
166
+ "label": "Person lying down",
167
+ "description": "A person is lying down"
168
+ }
169
+ ]
170
+ },
171
+ {
172
+ "id": "objects",
173
+ "name": "Vehicles & objects",
174
+ "checks": [
175
+ {
176
+ "id": "car",
177
+ "label": "Car",
178
+ "description": "A car is visible"
179
+ },
180
+ {
181
+ "id": "bus",
182
+ "label": "Bus",
183
+ "description": "A bus is visible"
184
+ },
185
+ {
186
+ "id": "bicycle",
187
+ "label": "Bicycle",
188
+ "description": "A bicycle is visible"
189
+ },
190
+ {
191
+ "id": "motorcycle",
192
+ "label": "Motorcycle",
193
+ "description": "A motorcycle is visible"
194
+ },
195
+ {
196
+ "id": "truck",
197
+ "label": "Truck",
198
+ "description": "A truck is visible"
199
+ },
200
+ {
201
+ "id": "taxi",
202
+ "label": "Taxi",
203
+ "description": "A marked taxi is visible"
204
+ },
205
+ {
206
+ "id": "boat",
207
+ "label": "Boat",
208
+ "description": "A boat is visible"
209
+ },
210
+ {
211
+ "id": "train",
212
+ "label": "Train",
213
+ "description": "A train is visible"
214
+ },
215
+ {
216
+ "id": "traffic_light",
217
+ "label": "Traffic light",
218
+ "description": "A traffic light is visible"
219
+ },
220
+ {
221
+ "id": "street_sign",
222
+ "label": "Street sign",
223
+ "description": "A street sign is visible"
224
+ },
225
+ {
226
+ "id": "billboard",
227
+ "label": "Billboard",
228
+ "description": "A large advertising billboard is visible"
229
+ },
230
+ {
231
+ "id": "bench",
232
+ "label": "Bench",
233
+ "description": "A bench is visible"
234
+ },
235
+ {
236
+ "id": "chair",
237
+ "label": "Chair",
238
+ "description": "A chair is visible"
239
+ },
240
+ {
241
+ "id": "table",
242
+ "label": "Table",
243
+ "description": "A table is visible"
244
+ },
245
+ {
246
+ "id": "trash_bin",
247
+ "label": "Trash bin",
248
+ "description": "A trash bin is visible"
249
+ },
250
+ {
251
+ "id": "traffic_cone",
252
+ "label": "Traffic cone",
253
+ "description": "A traffic cone is visible"
254
+ },
255
+ {
256
+ "id": "bollard",
257
+ "label": "Bollard",
258
+ "description": "A bollard or short street barrier post is visible"
259
+ },
260
+ {
261
+ "id": "fence",
262
+ "label": "Fence",
263
+ "description": "A fence is visible"
264
+ },
265
+ {
266
+ "id": "streetlamp",
267
+ "label": "Streetlamp",
268
+ "description": "A streetlamp is visible"
269
+ },
270
+ {
271
+ "id": "shop_window",
272
+ "label": "Shop window",
273
+ "description": "A shop display window is visible"
274
+ },
275
+ {
276
+ "id": "door",
277
+ "label": "Door",
278
+ "description": "A door is visible"
279
+ },
280
+ {
281
+ "id": "stairs",
282
+ "label": "Stairs",
283
+ "description": "Stairs are visible"
284
+ },
285
+ {
286
+ "id": "ramp",
287
+ "label": "Ramp",
288
+ "description": "A ramp is visible"
289
+ },
290
+ {
291
+ "id": "clock",
292
+ "label": "Clock",
293
+ "description": "A clock face is visible"
294
+ },
295
+ {
296
+ "id": "flag",
297
+ "label": "Flag",
298
+ "description": "A flag is visible"
299
+ },
300
+ {
301
+ "id": "food_stall",
302
+ "label": "Food stall",
303
+ "description": "A food stall or food cart is visible"
304
+ },
305
+ {
306
+ "id": "bottle",
307
+ "label": "Bottle",
308
+ "description": "A bottle is visible"
309
+ },
310
+ {
311
+ "id": "cup",
312
+ "label": "Cup",
313
+ "description": "A drinking cup is visible"
314
+ },
315
+ {
316
+ "id": "suitcase",
317
+ "label": "Suitcase",
318
+ "description": "A suitcase is visible"
319
+ },
320
+ {
321
+ "id": "screen",
322
+ "label": "Screen",
323
+ "description": "An electronic display screen is visible"
324
+ },
325
+ {
326
+ "id": "fire_hydrant",
327
+ "label": "Fire hydrant",
328
+ "description": "A fire hydrant is visible"
329
+ },
330
+ {
331
+ "id": "scaffolding",
332
+ "label": "Scaffolding",
333
+ "description": "Construction scaffolding is visible"
334
+ }
335
+ ]
336
+ },
337
+ {
338
+ "id": "nature",
339
+ "name": "Animals & environment",
340
+ "checks": [
341
+ {
342
+ "id": "dog",
343
+ "label": "Dog",
344
+ "description": "A real dog is visible"
345
+ },
346
+ {
347
+ "id": "cat",
348
+ "label": "Cat",
349
+ "description": "A real cat is visible"
350
+ },
351
+ {
352
+ "id": "bird",
353
+ "label": "Bird",
354
+ "description": "A real bird is visible"
355
+ },
356
+ {
357
+ "id": "horse",
358
+ "label": "Horse",
359
+ "description": "A real horse is visible"
360
+ },
361
+ {
362
+ "id": "cow",
363
+ "label": "Cow",
364
+ "description": "A real cow is visible"
365
+ },
366
+ {
367
+ "id": "sheep",
368
+ "label": "Sheep",
369
+ "description": "A real sheep is visible"
370
+ },
371
+ {
372
+ "id": "goat",
373
+ "label": "Goat",
374
+ "description": "A real goat is visible"
375
+ },
376
+ {
377
+ "id": "skunk",
378
+ "label": "Skunk",
379
+ "description": "A real skunk is visible"
380
+ },
381
+ {
382
+ "id": "rabbit",
383
+ "label": "Rabbit",
384
+ "description": "A real rabbit is visible"
385
+ },
386
+ {
387
+ "id": "squirrel",
388
+ "label": "Squirrel",
389
+ "description": "A real squirrel is visible"
390
+ },
391
+ {
392
+ "id": "duck",
393
+ "label": "Duck",
394
+ "description": "A real duck is visible"
395
+ },
396
+ {
397
+ "id": "fish",
398
+ "label": "Fish",
399
+ "description": "A real fish is visible"
400
+ },
401
+ {
402
+ "id": "tree",
403
+ "label": "Tree",
404
+ "description": "A tree is visible"
405
+ },
406
+ {
407
+ "id": "grass",
408
+ "label": "Grass",
409
+ "description": "Grass is visible"
410
+ },
411
+ {
412
+ "id": "flowers",
413
+ "label": "Flowers",
414
+ "description": "Flowers are visible"
415
+ },
416
+ {
417
+ "id": "potted_plant",
418
+ "label": "Potted plant",
419
+ "description": "A potted plant is visible"
420
+ },
421
+ {
422
+ "id": "bush",
423
+ "label": "Bush",
424
+ "description": "A bush is visible"
425
+ },
426
+ {
427
+ "id": "mountain",
428
+ "label": "Mountain",
429
+ "description": "A mountain is visible"
430
+ },
431
+ {
432
+ "id": "water",
433
+ "label": "Water",
434
+ "description": "An exposed body of water is visible"
435
+ },
436
+ {
437
+ "id": "beach",
438
+ "label": "Beach",
439
+ "description": "A sandy beach is visible"
440
+ },
441
+ {
442
+ "id": "snow",
443
+ "label": "Snow",
444
+ "description": "Snow is visible"
445
+ },
446
+ {
447
+ "id": "rain",
448
+ "label": "Rain",
449
+ "description": "Falling rain is visible"
450
+ },
451
+ {
452
+ "id": "clouds",
453
+ "label": "Clouds",
454
+ "description": "Clouds are visible"
455
+ },
456
+ {
457
+ "id": "blue_sky",
458
+ "label": "Blue sky",
459
+ "description": "Blue sky is visible"
460
+ },
461
+ {
462
+ "id": "sun",
463
+ "label": "Sun",
464
+ "description": "The sun itself is visible"
465
+ },
466
+ {
467
+ "id": "moon",
468
+ "label": "Moon",
469
+ "description": "The moon itself is visible"
470
+ },
471
+ {
472
+ "id": "smoke",
473
+ "label": "Smoke",
474
+ "description": "Smoke is visible"
475
+ },
476
+ {
477
+ "id": "fire",
478
+ "label": "Fire",
479
+ "description": "Flames are visible"
480
+ },
481
+ {
482
+ "id": "rocks",
483
+ "label": "Rocks",
484
+ "description": "Natural rocks are visible"
485
+ },
486
+ {
487
+ "id": "fallen_leaves",
488
+ "label": "Fallen leaves",
489
+ "description": "Fallen leaves are visible on the ground"
490
+ },
491
+ {
492
+ "id": "puddle",
493
+ "label": "Puddle",
494
+ "description": "A puddle is visible"
495
+ },
496
+ {
497
+ "id": "animal_group",
498
+ "label": "Animal group",
499
+ "description": "Two or more real animals are visible"
500
+ }
501
+ ]
502
+ },
503
+ {
504
+ "id": "scene",
505
+ "name": "Scene & composition",
506
+ "checks": [
507
+ {
508
+ "id": "outdoors",
509
+ "label": "Outdoors",
510
+ "description": "The scene is outdoors"
511
+ },
512
+ {
513
+ "id": "indoors",
514
+ "label": "Indoors",
515
+ "description": "The scene is indoors"
516
+ },
517
+ {
518
+ "id": "street",
519
+ "label": "Street",
520
+ "description": "A street or road is visible"
521
+ },
522
+ {
523
+ "id": "sidewalk",
524
+ "label": "Sidewalk",
525
+ "description": "A sidewalk is visible"
526
+ },
527
+ {
528
+ "id": "crosswalk",
529
+ "label": "Crosswalk",
530
+ "description": "A marked pedestrian crossing is visible"
531
+ },
532
+ {
533
+ "id": "buildings",
534
+ "label": "Buildings",
535
+ "description": "Buildings are visible"
536
+ },
537
+ {
538
+ "id": "high_rise",
539
+ "label": "High rise",
540
+ "description": "A high-rise building is visible"
541
+ },
542
+ {
543
+ "id": "storefront",
544
+ "label": "Storefront",
545
+ "description": "A storefront is visible"
546
+ },
547
+ {
548
+ "id": "park",
549
+ "label": "Park",
550
+ "description": "The setting is visibly a park"
551
+ },
552
+ {
553
+ "id": "kitchen",
554
+ "label": "Kitchen",
555
+ "description": "The setting is visibly a kitchen"
556
+ },
557
+ {
558
+ "id": "office",
559
+ "label": "Office",
560
+ "description": "The setting is visibly an office"
561
+ },
562
+ {
563
+ "id": "living_room",
564
+ "label": "Living room",
565
+ "description": "The setting is visibly a living room"
566
+ },
567
+ {
568
+ "id": "daylight",
569
+ "label": "Daylight",
570
+ "description": "The scene is lit by daylight"
571
+ },
572
+ {
573
+ "id": "night",
574
+ "label": "Night",
575
+ "description": "The scene is visibly at night"
576
+ },
577
+ {
578
+ "id": "artificial_lighting",
579
+ "label": "Artificial lighting",
580
+ "description": "Artificial lights are visibly illuminating the scene"
581
+ },
582
+ {
583
+ "id": "shadows",
584
+ "label": "Shadows",
585
+ "description": "Distinct cast shadows are visible"
586
+ },
587
+ {
588
+ "id": "reflections",
589
+ "label": "Reflections",
590
+ "description": "Reflections are visible"
591
+ },
592
+ {
593
+ "id": "wet_ground",
594
+ "label": "Wet ground",
595
+ "description": "The ground appears wet"
596
+ },
597
+ {
598
+ "id": "visible_text",
599
+ "label": "Visible text",
600
+ "description": "Readable or recognizable text is visible"
601
+ },
602
+ {
603
+ "id": "advertising",
604
+ "label": "Advertising",
605
+ "description": "Advertising is visible"
606
+ },
607
+ {
608
+ "id": "road_markings",
609
+ "label": "Road markings",
610
+ "description": "Painted road markings are visible"
611
+ },
612
+ {
613
+ "id": "red_dominant_area",
614
+ "label": "Red dominant area",
615
+ "description": "A large red area is visible"
616
+ },
617
+ {
618
+ "id": "blue_dominant_area",
619
+ "label": "Blue dominant area",
620
+ "description": "A large blue area is visible"
621
+ },
622
+ {
623
+ "id": "green_dominant_area",
624
+ "label": "Green dominant area",
625
+ "description": "A large green area is visible"
626
+ },
627
+ {
628
+ "id": "yellow_dominant_area",
629
+ "label": "Yellow dominant area",
630
+ "description": "A large yellow area is visible"
631
+ },
632
+ {
633
+ "id": "closeup",
634
+ "label": "Closeup",
635
+ "description": "The framing is a close-up view"
636
+ },
637
+ {
638
+ "id": "wide_view",
639
+ "label": "Wide view",
640
+ "description": "The framing shows a wide scene"
641
+ },
642
+ {
643
+ "id": "blur",
644
+ "label": "Blur",
645
+ "description": "Substantial image blur is visible"
646
+ },
647
+ {
648
+ "id": "occlusion",
649
+ "label": "Occlusion",
650
+ "description": "A main subject is partly hidden by another object"
651
+ },
652
+ {
653
+ "id": "dense_scene",
654
+ "label": "Dense scene",
655
+ "description": "Many distinct objects fill the scene"
656
+ },
657
+ {
658
+ "id": "clear_foreground",
659
+ "label": "Clear foreground",
660
+ "description": "A clear foreground subject is visible"
661
+ },
662
+ {
663
+ "id": "distant_background",
664
+ "label": "Distant background",
665
+ "description": "A distant background is visible"
666
+ }
667
+ ]
668
+ }
669
+ ],
670
+ "sample": {
671
+ "url": "/static/sample-street.jpg",
672
+ "name": "Times Square",
673
+ "author": "ISO Legacy",
674
+ "source": "https://commons.wikimedia.org/wiki/File:Times_Square_(New_York_City).jpg",
675
+ "license": "CC0 1.0"
676
+ }
677
+ }
gemma_rlcd/web.py CHANGED
@@ -11,9 +11,11 @@ from concurrent.futures import ThreadPoolExecutor
11
  from contextlib import asynccontextmanager
12
  from pathlib import Path
13
  from tempfile import TemporaryDirectory
 
14
 
 
15
  from fastapi import FastAPI, Request
16
- from fastapi.responses import FileResponse, JSONResponse
17
  from fastapi.staticfiles import StaticFiles
18
  from PIL import Image, ImageOps, UnidentifiedImageError
19
  from starlette.datastructures import UploadFile
@@ -104,8 +106,8 @@ def read_spec(encoded: str) -> tuple[dict, dict]:
104
  field_count = sum(
105
  len(q.criteria) if isinstance(q, Independent) else 1 for q in questions.values()
106
  )
107
- if field_count > 64:
108
- raise ValueError("At most 64 individual fields or independent labels can run together")
109
  media = spec.get("media", [])
110
  if not isinstance(media, list) or len(media) > 10:
111
  raise ValueError("Attach at most 8 images, one audio clip, and one video")
@@ -121,6 +123,32 @@ def read_spec(encoded: str) -> tuple[dict, dict]:
121
  return spec, questions
122
 
123
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
124
  def inspect_media(path: Path) -> dict:
125
  try:
126
  process = subprocess.run(
@@ -240,6 +268,7 @@ class Runtime:
240
  self.pool = ThreadPoolExecutor(max_workers=1, thread_name_prefix="decision-model")
241
  self.lock = asyncio.Lock()
242
  self.backend = None
 
243
  self.error = None
244
  self.load_seconds = None
245
 
@@ -249,9 +278,13 @@ class Runtime:
249
  if self.factory is None:
250
  from .json_backend import JSONMLXBackend
251
 
252
- self.backend = JSONMLXBackend(self.model, branch_batch_size=8)
253
  else:
254
- self.backend = self.factory()
 
 
 
 
255
  self.load_seconds = time.perf_counter() - started
256
 
257
  try:
@@ -261,7 +294,7 @@ class Runtime:
261
  self.error = "The model could not load. Check the model path and server log."
262
 
263
  def evaluate(
264
- self, spec: dict, questions: dict, paths: list[Path], comparison: bool = False
265
  ) -> dict:
266
  started = time.perf_counter()
267
  media, metadata = prepare_media(paths, spec.get("media", []))
@@ -271,7 +304,19 @@ class Runtime:
271
  if comparison:
272
  from .comparison import compare
273
 
274
- result, details = compare(self.backend, state, questions, normalized - started)
 
 
 
 
 
 
 
 
 
 
 
 
275
  else:
276
  result = DecisionEngine(self.backend).system_one(state, questions)
277
  finished = time.perf_counter()
@@ -332,6 +377,90 @@ def create_app(model: str, work_dir: Path, backend_factory=None) -> FastAPI:
332
  async def index():
333
  return FileResponse(STATIC / "index.html")
334
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
335
  @app.get("/api/status")
336
  async def status():
337
  return {
@@ -359,46 +488,21 @@ def create_app(model: str, work_dir: Path, backend_factory=None) -> FastAPI:
359
  async with runtime.lock:
360
  started = time.perf_counter()
361
  try:
362
- async with request.form(
363
- max_files=10, max_fields=1, max_part_size=1024 * 1024
364
- ) as form:
365
- if set(form) - {"spec", "media"} or not isinstance(form.get("spec"), str):
366
- raise ValueError("Submit a JSON spec and optional media attachments")
367
- spec, questions = read_spec(form["spec"])
368
- files = form.getlist("media")
369
- if len(files) != len(spec.get("media", [])) or not all(
370
- isinstance(file, UploadFile) for file in files
371
- ):
372
- raise ValueError("Attachment files do not match the request")
373
- with TemporaryDirectory(prefix="run-", dir=work_dir) as directory:
374
- paths, total = [], 0
375
- for index, file in enumerate(files):
376
- suffix = Path(file.filename or "upload").suffix.lower()
377
- if len(suffix) > 10 or not suffix.replace(".", "").isalnum():
378
- suffix = ".bin"
379
- path = Path(directory) / f"attachment-{index}{suffix}"
380
- with path.open("wb") as output:
381
- while chunk := await file.read(1024 * 1024):
382
- total += len(chunk)
383
- if total > MAX_UPLOAD_BYTES:
384
- raise ValueError(
385
- "Attachments exceed the 200 MB total upload limit"
386
- )
387
- output.write(chunk)
388
- paths.append(path)
389
- future = asyncio.get_running_loop().run_in_executor(
390
- runtime.pool,
391
- runtime.evaluate,
392
- spec,
393
- questions,
394
- paths,
395
- request.url.path == "/api/compare",
396
- )
397
- try:
398
- result = await asyncio.shield(future)
399
- except asyncio.CancelledError:
400
- await future
401
- raise
402
  result["request_seconds"] = time.perf_counter() - started
403
  return JSONResponse(result)
404
  except (ValueError, TypeError, KeyError) as exc:
 
11
  from contextlib import asynccontextmanager
12
  from pathlib import Path
13
  from tempfile import TemporaryDirectory
14
+ from threading import Event
15
 
16
+ from anyio import CancelScope
17
  from fastapi import FastAPI, Request
18
+ from fastapi.responses import FileResponse, JSONResponse, StreamingResponse
19
  from fastapi.staticfiles import StaticFiles
20
  from PIL import Image, ImageOps, UnidentifiedImageError
21
  from starlette.datastructures import UploadFile
 
106
  field_count = sum(
107
  len(q.criteria) if isinstance(q, Independent) else 1 for q in questions.values()
108
  )
109
+ if field_count > 128:
110
+ raise ValueError("At most 128 individual fields or independent labels can run together")
111
  media = spec.get("media", [])
112
  if not isinstance(media, list) or len(media) > 10:
113
  raise ValueError("Attach at most 8 images, one audio clip, and one video")
 
123
  return spec, questions
124
 
125
 
126
+ async def read_upload(request: Request, directory: str):
127
+ async with request.form(max_files=10, max_fields=1, max_part_size=1024 * 1024) as form:
128
+ if set(form) - {"spec", "media"} or not isinstance(form.get("spec"), str):
129
+ raise ValueError("Submit a JSON spec and optional media attachments")
130
+ spec, questions = read_spec(form["spec"])
131
+ files = form.getlist("media")
132
+ if len(files) != len(spec.get("media", [])) or not all(
133
+ isinstance(file, UploadFile) for file in files
134
+ ):
135
+ raise ValueError("Attachment files do not match the request")
136
+ paths, total = [], 0
137
+ for index, file in enumerate(files):
138
+ suffix = Path(file.filename or "upload").suffix.lower()
139
+ if len(suffix) > 10 or not suffix.replace(".", "").isalnum():
140
+ suffix = ".bin"
141
+ path = Path(directory) / f"attachment-{index}{suffix}"
142
+ with path.open("wb") as output:
143
+ while chunk := await file.read(1024 * 1024):
144
+ total += len(chunk)
145
+ if total > MAX_UPLOAD_BYTES:
146
+ raise ValueError("Attachments exceed the 200 MB total upload limit")
147
+ output.write(chunk)
148
+ paths.append(path)
149
+ return spec, questions, paths
150
+
151
+
152
  def inspect_media(path: Path) -> dict:
153
  try:
154
  process = subprocess.run(
 
268
  self.pool = ThreadPoolExecutor(max_workers=1, thread_name_prefix="decision-model")
269
  self.lock = asyncio.Lock()
270
  self.backend = None
271
+ self.normal_backend = None
272
  self.error = None
273
  self.load_seconds = None
274
 
 
278
  if self.factory is None:
279
  from .json_backend import JSONMLXBackend
280
 
281
+ backend = JSONMLXBackend(self.model, branch_batch_size=8)
282
  else:
283
+ backend = self.factory()
284
+ from .comparison import generation_backend
285
+
286
+ self.normal_backend = generation_backend(backend)
287
+ self.backend = backend
288
  self.load_seconds = time.perf_counter() - started
289
 
290
  try:
 
294
  self.error = "The model could not load. Check the model path and server log."
295
 
296
  def evaluate(
297
+ self, spec: dict, questions: dict, paths: list[Path], comparison: bool = False, emit=None
298
  ) -> dict:
299
  started = time.perf_counter()
300
  media, metadata = prepare_media(paths, spec.get("media", []))
 
304
  if comparison:
305
  from .comparison import compare
306
 
307
+ result, details = (
308
+ compare(
309
+ self.backend,
310
+ state,
311
+ questions,
312
+ normalized - started,
313
+ emit=emit,
314
+ concurrent=True,
315
+ normal_backend=self.normal_backend,
316
+ )
317
+ if emit
318
+ else compare(self.backend, state, questions, normalized - started)
319
+ )
320
  else:
321
  result = DecisionEngine(self.backend).system_one(state, questions)
322
  finished = time.perf_counter()
 
377
  async def index():
378
  return FileResponse(STATIC / "index.html")
379
 
380
+ @app.get("/demo")
381
+ async def demo():
382
+ return FileResponse(STATIC / "demo.html")
383
+
384
+ @app.post("/api/compare-stream")
385
+ async def stream(request: Request):
386
+ if runtime.backend is None or runtime.error:
387
+ return JSONResponse(
388
+ {"error": runtime.error or "The model is still loading."}, status_code=503
389
+ )
390
+ if runtime.lock.locked():
391
+ return JSONResponse(
392
+ {"error": "Another evaluation is running. Wait for it to finish."}, status_code=429
393
+ )
394
+
395
+ await runtime.lock.acquire()
396
+ try:
397
+ directory = TemporaryDirectory(prefix="stream-", dir=work_dir)
398
+ except BaseException:
399
+ runtime.lock.release()
400
+ raise
401
+ try:
402
+ spec, questions, paths = await read_upload(request, directory.name)
403
+ except BaseException as exc:
404
+ directory.cleanup()
405
+ runtime.lock.release()
406
+ if isinstance(exc, (ValueError, TypeError, KeyError)):
407
+ return JSONResponse({"error": str(exc)}, status_code=400)
408
+ raise
409
+
410
+ async def events():
411
+ loop = asyncio.get_running_loop()
412
+ queue = asyncio.Queue()
413
+ stopped = Event()
414
+
415
+ class StreamStopped(Exception):
416
+ pass
417
+
418
+ def emit(event):
419
+ if stopped.is_set():
420
+ raise StreamStopped()
421
+ loop.call_soon_threadsafe(queue.put_nowait, event)
422
+
423
+ def evaluate():
424
+ try:
425
+ result = runtime.evaluate(spec, questions, paths, True, emit)
426
+ emit({"type": "complete", "result": result})
427
+ except StreamStopped:
428
+ pass
429
+ except (ValueError, TypeError, KeyError) as exc:
430
+ loop.call_soon_threadsafe(
431
+ queue.put_nowait, {"type": "error", "error": str(exc)}
432
+ )
433
+ except Exception:
434
+ LOGGER.exception("Streaming evaluation failed")
435
+ if not stopped.is_set():
436
+ loop.call_soon_threadsafe(
437
+ queue.put_nowait,
438
+ {
439
+ "type": "error",
440
+ "error": "Evaluation failed. Check the media or reduce the input size. Details are in the server log.",
441
+ },
442
+ )
443
+
444
+ future = loop.run_in_executor(runtime.pool, evaluate)
445
+ future.add_done_callback(lambda _: queue.put_nowait(None))
446
+ try:
447
+ yield json.dumps({"type": "accepted"}) + "\n"
448
+ while (event := await queue.get()) is not None:
449
+ yield json.dumps(event, allow_nan=False) + "\n"
450
+ finally:
451
+ stopped.set()
452
+ # Keep uploads and the GPU lock alive until the worker stops.
453
+ with CancelScope(shield=True):
454
+ try:
455
+ await asyncio.shield(future)
456
+ finally:
457
+ directory.cleanup()
458
+ runtime.lock.release()
459
+
460
+ return StreamingResponse(
461
+ events(), media_type="application/x-ndjson", headers={"X-Accel-Buffering": "no"}
462
+ )
463
+
464
  @app.get("/api/status")
465
  async def status():
466
  return {
 
488
  async with runtime.lock:
489
  started = time.perf_counter()
490
  try:
491
+ with TemporaryDirectory(prefix="run-", dir=work_dir) as directory:
492
+ spec, questions, paths = await read_upload(request, directory)
493
+ future = asyncio.get_running_loop().run_in_executor(
494
+ runtime.pool,
495
+ runtime.evaluate,
496
+ spec,
497
+ questions,
498
+ paths,
499
+ request.url.path == "/api/compare",
500
+ )
501
+ try:
502
+ result = await asyncio.shield(future)
503
+ except asyncio.CancelledError:
504
+ await future
505
+ raise
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
506
  result["request_seconds"] = time.perf_counter() - started
507
  return JSONResponse(result)
508
  except (ValueError, TypeError, KeyError) as exc:
tests/test_comparison.py CHANGED
@@ -3,7 +3,7 @@ import json
3
  import pytest
4
 
5
  from gemma_rlcd.comparison import discrete_answers, generation_task, parse_generated
6
- from gemma_rlcd.core import Choice, Independent, Noul, Score
7
 
8
 
9
  def questions():
@@ -58,3 +58,117 @@ def test_agreement_uses_grade_mode_and_boolean_thresholds():
58
  "presence": {"type": "independent", "probabilities": {"car": 0.9, "person": 0.1}},
59
  }
60
  ) == {"grade": 2, "moving": True, "presence": {"car": True, "person": False}}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3
  import pytest
4
 
5
  from gemma_rlcd.comparison import discrete_answers, generation_task, parse_generated
6
+ from gemma_rlcd.core import Choice, Independent, Noul, Score, State
7
 
8
 
9
  def questions():
 
58
  "presence": {"type": "independent", "probabilities": {"car": 0.9, "person": 0.1}},
59
  }
60
  ) == {"grade": 2, "moving": True, "presence": {"car": True, "person": False}}
61
+
62
+
63
+ def test_streaming_generation_uses_real_chunks_and_validates_complete_answer(monkeypatch):
64
+ import sys
65
+ from types import SimpleNamespace
66
+
67
+ from gemma_rlcd import comparison
68
+
69
+ seen = []
70
+ parts = ['{"visible":', "true", "}"]
71
+
72
+ def stream(*args, **kwargs):
73
+ for index, text in enumerate(parts):
74
+ assert len(seen) == index
75
+ yield SimpleNamespace(
76
+ text=text,
77
+ generation_tokens=index + 1,
78
+ prompt_tokens=50,
79
+ finish_reason="stop" if index == len(parts) - 1 else None,
80
+ )
81
+
82
+ monkeypatch.setitem(
83
+ sys.modules, "mlx_vlm", SimpleNamespace(generate=None, stream_generate=stream)
84
+ )
85
+ monkeypatch.setattr(comparison, "prepare_generation", lambda *args: ("prompt", {}))
86
+ backend = SimpleNamespace(
87
+ model=None,
88
+ processor=None,
89
+ tokenizer=SimpleNamespace(encode=lambda *args, **kwargs: [1]),
90
+ mx=SimpleNamespace(synchronize=lambda: None),
91
+ )
92
+ result = comparison.generate_answers(
93
+ backend,
94
+ State(text="test"),
95
+ {"visible": Noul("Visible?")},
96
+ on_token=lambda text, count: seen.append((text, count)),
97
+ )
98
+ assert seen == list(zip(parts, [1, 2, 3], strict=True))
99
+ assert result["valid"] and result["answers"] == {"visible": True}
100
+ assert result["raw_text"] == "".join(parts)
101
+
102
+
103
+ def test_concurrent_comparison_streams_both_paths_before_either_finishes(monkeypatch):
104
+ from threading import Event
105
+ from types import SimpleNamespace
106
+
107
+ from gemma_rlcd import comparison
108
+ from gemma_rlcd.core import TokenScores
109
+
110
+ scored, generated = Event(), Event()
111
+ events = []
112
+
113
+ class Backend:
114
+ last_stats = {}
115
+ processor = SimpleNamespace(tokenizer=SimpleNamespace(mutable=[]))
116
+
117
+ def symbols(self, count):
118
+ return ("A", "B")
119
+
120
+ def score_questions(self, state, questions, on_scores):
121
+ assert generated.wait(2), "Generation never started alongside scoring"
122
+ result = TokenScores((5, 0), 1, 20)
123
+ on_scores([(0, result)])
124
+ scored.set()
125
+ return [result]
126
+
127
+ backend = Backend()
128
+
129
+ def generate(worker, state, questions, on_token):
130
+ assert worker.processor is not backend.processor
131
+ assert worker.processor.tokenizer is not backend.processor.tokenizer
132
+ on_token('{"visible":', 1)
133
+ generated.set()
134
+ assert scored.wait(2), "Scoring never completed while generation was active"
135
+ on_token("true}", 2)
136
+ return {"answers": {"visible": True}, "valid": True, "inference_seconds": 0}
137
+
138
+ monkeypatch.setattr(comparison, "generate_answers", generate)
139
+ _, result = comparison.compare(
140
+ backend,
141
+ State(text="scene"),
142
+ {"visible": Noul("Visible?")},
143
+ 0,
144
+ emit=events.append,
145
+ concurrent=True,
146
+ )
147
+ assert events[0]["type"] == "race_start"
148
+ first_token = next(i for i, event in enumerate(events) if event["type"] == "token")
149
+ first_answer = next(i for i, event in enumerate(events) if event["type"] == "answer")
150
+ first_finish = next(i for i, event in enumerate(events) if event["type"] == "phase_complete")
151
+ assert first_token < first_answer < first_finish
152
+ assert result["agreement"] == {"visible": True}
153
+ assert result["methodology"]["execution"] == "concurrent_shared_gpu"
154
+
155
+
156
+ @pytest.mark.parametrize("cancel_at", ["race_start", "phase_start"])
157
+ def test_concurrent_comparison_can_stop_before_or_after_worker_launch(cancel_at):
158
+ from types import SimpleNamespace
159
+
160
+ from gemma_rlcd.comparison import compare
161
+
162
+ def emit(event):
163
+ if event["type"] == cancel_at:
164
+ raise RuntimeError("Client stopped")
165
+
166
+ with pytest.raises(RuntimeError, match="Client stopped"):
167
+ compare(
168
+ SimpleNamespace(),
169
+ State(text="scene"),
170
+ {"visible": Noul("Visible?")},
171
+ 0,
172
+ emit=emit,
173
+ concurrent=True,
174
+ )
tests/test_decisions.py CHANGED
@@ -178,3 +178,25 @@ def test_batch_result_count_is_checked():
178
 
179
  with pytest.raises(ValueError, match="wrong number of question results"):
180
  DecisionEngine(Broken([])).system_one(State(text="test"), {"test": Noul("True?")})
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
178
 
179
  with pytest.raises(ValueError, match="wrong number of question results"):
180
  DecisionEngine(Broken([])).system_one(State(text="test"), {"test": Noul("True?")})
181
+
182
+
183
+ def test_streamed_decisions_arrive_before_later_batches_and_match_final_answers():
184
+ seen = []
185
+
186
+ class Streaming(StubBackend):
187
+ def score_questions(self, state, questions, on_scores=None):
188
+ scores = [TokenScores((2.0, 0.0), 0.8, 50), TokenScores((0.0, 2.0), 0.8, 50)]
189
+ on_scores([(0, scores[0])])
190
+ assert len(seen) == 1
191
+ on_scores([(1, scores[1])])
192
+ return scores
193
+
194
+ result = DecisionEngine(Streaming([])).system_one(
195
+ State(text="evidence"),
196
+ {"visible": Independent("Which?", {"cat": "Cat", "dog": "Dog"})},
197
+ on_answer=lambda path, answer: seen.append((path, answer)),
198
+ )
199
+ assert [path for path, _ in seen] == [("visible", "cat"), ("visible", "dog")]
200
+ assert result["answers"]["visible"]["probabilities"] == {
201
+ path[1]: answer["probabilities"]["yes"] for path, answer in seen
202
+ }
tests/test_demo.cjs ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ const assert = require('node:assert/strict');
2
+ const {test} = require('node:test');
3
+ const demo = require('../gemma_rlcd/static/demo-utils.js');
4
+ const config = require('../gemma_rlcd/static/visual-demo.json');
5
+ test('visual schemas contain distinct meaningful checks at each size', () => {
6
+ for (const size of [32,64,128]) {
7
+ const questions = demo.questions(config,size);
8
+ assert.equal(Object.values(questions).reduce((n,q) => n + Object.keys(q.criteria).length,0),size);
9
+ assert.equal(new Set(Object.values(questions).flatMap(q => Object.values(q.criteria))).size,size);
10
+ }
11
+ assert.throws(() => demo.questions(config,129));
12
+ });
13
+ test('partial JSON reveals only complete booleans with a delimiter', () => {
14
+ const input = '{"people":{"person":true,"walking":false},"objects":{"car":true}}';
15
+ assert.deepEqual({...demo.partialBooleans(input)}, {'people.person':true,'people.walking':false,'objects.car':true});
16
+ const prefix = '{"people":{"person":true';
17
+ assert.deepEqual({...demo.partialBooleans(prefix)}, {});
18
+ assert.deepEqual({...demo.partialBooleans(prefix + ',"walking":fa')}, {'people.person':true});
19
+ assert.deepEqual({...demo.partialBooleans('{"people":{"person":"true"}')}, {});
20
+ assert.deepEqual({...demo.partialBooleans('{"people":{"person":truefake,')}, {});
21
+ assert.deepEqual({...demo.partialBooleans('```json\n' + input + '\n```')}, {...demo.partialBooleans(input)});
22
+ });
tests/test_json_backend.py CHANGED
@@ -27,9 +27,11 @@ def test_complete_candidate_likelihood_uses_suffix_and_ignores_padding(batch_siz
27
  )
28
  )
29
  fields = [(0, FieldTokens((9,), ((1, 2), (1, 3), (4,))))]
30
- results, batches = backend._sequence_scores([], 100, fields)
 
31
  likelihoods = [math.exp(value) for value in results[0].logits]
32
  assert likelihoods == pytest.approx([0.09, 0.81, 0.1], abs=1e-6)
33
  assert results[0].allowed_token_mass == pytest.approx(1, abs=1e-6)
34
  assert sum(batches) == 3
35
  assert max(batches) <= batch_size
 
 
27
  )
28
  )
29
  fields = [(0, FieldTokens((9,), ((1, 2), (1, 3), (4,))))]
30
+ streamed = []
31
+ results, batches = backend._sequence_scores([], 100, fields, on_scores=streamed.append)
32
  likelihoods = [math.exp(value) for value in results[0].logits]
33
  assert likelihoods == pytest.approx([0.09, 0.81, 0.1], abs=1e-6)
34
  assert results[0].allowed_token_mass == pytest.approx(1, abs=1e-6)
35
  assert sum(batches) == 3
36
  assert max(batches) <= batch_size
37
+ assert streamed == [[(0, results[0])]]
tests/test_web.py CHANGED
@@ -264,10 +264,10 @@ def test_independent_expansion_is_bounded():
264
  "labels": {
265
  "type": "independent",
266
  "instructions": "Check each",
267
- "criteria": {str(i): "An option" for i in range(65)},
268
  }
269
  }
270
- with pytest.raises(ValueError, match="64 individual fields"):
271
  read_spec(json.dumps(spec))
272
 
273
 
@@ -442,3 +442,78 @@ def test_invalid_generated_answer_remains_visible_without_speedup_claim(playgrou
442
  assert data["comparison"]["agreement"]["animal"] is None
443
  assert data["comparison"]["normal"]["raw_text"] == "not json"
444
  assert data["answers"]["animal"]["choice"] == "cat"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
264
  "labels": {
265
  "type": "independent",
266
  "instructions": "Check each",
267
+ "criteria": {str(i): "An option" for i in range(129)},
268
  }
269
  }
270
+ with pytest.raises(ValueError, match="128 individual fields"):
271
  read_spec(json.dumps(spec))
272
 
273
 
 
442
  assert data["comparison"]["agreement"]["animal"] is None
443
  assert data["comparison"]["normal"]["raw_text"] == "not json"
444
  assert data["answers"]["animal"]["choice"] == "cat"
445
+
446
+
447
+ def test_visual_demo_has_128_distinct_checks_and_supports_each_size(playground):
448
+ client, backend, _ = playground
449
+ assert "One scene." in client.get("/demo").text
450
+ config = client.get("/static/visual-demo.json").json()
451
+ assert len(config["groups"]) == 4
452
+ for size in (32, 64, 128):
453
+ spec = {
454
+ "text": "scene",
455
+ "instructions": config["instructions"],
456
+ "questions": {
457
+ group["id"]: {
458
+ "type": "independent",
459
+ "instructions": "Which checks are visible?",
460
+ "criteria": {
461
+ check["id"]: check["description"] for check in group["checks"][: size // 4]
462
+ },
463
+ }
464
+ for group in config["groups"]
465
+ },
466
+ }
467
+ response = client.post("/api/run", data={"spec": json.dumps(spec)})
468
+ assert response.status_code == 200, response.text
469
+ assert len(backend.calls[-1][1]) == size
470
+
471
+
472
+ def test_streaming_comparison_orders_events_and_cleans_uploads(playground, monkeypatch):
473
+ from gemma_rlcd import comparison
474
+
475
+ client, _, directory = playground
476
+
477
+ def generate(backend, state, questions, on_token=None):
478
+ on_token('{"animal":', 1)
479
+ on_token('"cat"}', 2)
480
+ return {
481
+ "answers": {"animal": "cat"},
482
+ "raw_text": '{"animal":"cat"}',
483
+ "valid": True,
484
+ "error": None,
485
+ "inference_seconds": 1,
486
+ "output_tokens": 2,
487
+ }
488
+
489
+ monkeypatch.setattr(comparison, "generate_answers", generate)
490
+ response = client.post("/api/compare-stream", data={"spec": json.dumps(request_spec())})
491
+ assert response.status_code == 200
492
+ assert response.headers["content-type"].startswith("application/x-ndjson")
493
+ events = [json.loads(line) for line in response.text.splitlines()]
494
+ assert [event["type"] for event in events[:2]] == ["accepted", "race_start"]
495
+ assert events[-1]["type"] == "complete"
496
+ assert [event["type"] for event in events if event.get("method") == "parallel"] == [
497
+ "phase_start",
498
+ "answer",
499
+ "phase_complete",
500
+ ]
501
+ assert [event["type"] for event in events if event.get("method") == "normal"] == [
502
+ "phase_start",
503
+ "token",
504
+ "token",
505
+ "phase_complete",
506
+ ]
507
+ assert next(event for event in events if event["type"] == "answer")["value"] == "cat"
508
+ assert events[-1]["result"]["comparison"]["normal"]["valid"]
509
+ assert events[-1]["result"]["comparison"]["methodology"]["execution"] == "concurrent_shared_gpu"
510
+ assert not client.get("/api/status").json()["busy"]
511
+ assert list(directory.iterdir()) == []
512
+
513
+
514
+ def test_stream_validation_failure_releases_gpu_lock(playground):
515
+ client, _, directory = playground
516
+ response = client.post("/api/compare-stream", data={"spec": "{}"})
517
+ assert response.status_code == 400
518
+ assert not client.get("/api/status").json()["busy"]
519
+ assert list(directory.iterdir()) == []