Any-to-Any
MLX
Safetensors
gemma4
mlx-vlm
rlcd
multimodal
classification
parallel-inference
image-text-to-text
audio
video
4-bit precision
Instructions to use larkooo/gemma-e2b-rlcd with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use larkooo/gemma-e2b-rlcd with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] huggingface-cli download --local-dir gemma-e2b-rlcd larkooo/gemma-e2b-rlcd
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
Download tests/test_web.py from larkooo/gemma-e2b-rlcd: direct link, hf CLI and curl.
- Browser
- Download file 18.7 kB
-
https://huggingface.co/larkooo/gemma-e2b-rlcd/resolve/main/tests/test_web.py
- Command line
-
hf download hf://larkooo/gemma-e2b-rlcd/tests/test_web.py
-
curl -L -o test_web.py https://huggingface.co/larkooo/gemma-e2b-rlcd/resolve/main/tests/test_web.py
18.7 kB
| import io | |
| import json | |
| import shutil | |
| import subprocess | |
| import time | |
| import wave | |
| from concurrent.futures import ThreadPoolExecutor | |
| from pathlib import Path | |
| from threading import Event | |
| import pytest | |
| from fastapi.testclient import TestClient | |
| from PIL import Image | |
| from gemma_rlcd.comparison import generation_task | |
| from gemma_rlcd.core import TokenScores, parse_question | |
| from gemma_rlcd.web import create_app, read_spec | |
| class FakeBackend: | |
| def __init__(self): | |
| self.last_stats = {"branch_batch_sizes": [1], "prefix_prefills": 1} | |
| self.calls = [] | |
| def symbols(self, count): | |
| return tuple(chr(65 + i) for i in range(count)) | |
| def score_batch(self, state, requests): | |
| self.calls.append((state, requests)) | |
| for path in (*state.images, *state.audio, *state.videos): | |
| from pathlib import Path | |
| assert Path(path).is_file() | |
| return [ | |
| TokenScores(tuple(float(-i) for i in range(len(request.symbols))), 1.0, 123) | |
| for request in requests | |
| ] | |
| def playground(tmp_path): | |
| backend = FakeBackend() | |
| app = create_app("test-model", tmp_path, backend_factory=lambda: backend) | |
| with TestClient(app, base_url="http://localhost") as client: | |
| for _ in range(100): | |
| if client.get("/api/status").json()["ready"]: | |
| break | |
| time.sleep(0.01) | |
| assert client.get("/api/status").json()["ready"] | |
| yield client, backend, tmp_path | |
| def request_spec(): | |
| return { | |
| "text": "A cat", | |
| "questions": { | |
| "animal": { | |
| "type": "choice", | |
| "instructions": "Which animal?", | |
| "criteria": {"cat": "A cat", "dog": "A dog"}, | |
| } | |
| }, | |
| } | |
| def test_large_demo_presets_are_served_and_all_decisions_reach_backend(playground): | |
| client, backend, _ = playground | |
| response = client.get("/static/demo-presets.json") | |
| assert response.status_code == 200 | |
| presets = response.json() | |
| for name, count in [ | |
| ("support_matrix", 28), | |
| ("inbox_matrix", 32), | |
| ("ticket_flags", 32), | |
| ("catalog_choices", 4), | |
| ]: | |
| spec = presets[name] | |
| response = client.post("/api/run", data={"spec": json.dumps(spec)}) | |
| assert response.status_code == 200 | |
| assert len(backend.calls[-1][1]) == count | |
| assert backend.calls[-1][0].text == spec["text"] | |
| assert set(response.json()["answers"]) == set(spec["questions"]) | |
| def test_demo_presets_preserve_the_benchmark_prompt_after_editor_serialization(playground): | |
| client, _, _ = playground | |
| presets = client.get("/static/demo-presets.json").json() | |
| cases = json.loads((Path(__file__).parents[1] / "examples/demo-workloads.json").read_text()) | |
| cases = {case["name"]: case for case in cases["cases"]} | |
| for preset_name, case_name in [ | |
| ("support_matrix", "support_28"), | |
| ("inbox_matrix", "inbox_32"), | |
| ("ticket_flags", "ticket_flags_32"), | |
| ("catalog_choices", "catalog_64"), | |
| ]: | |
| preset = presets[preset_name] | |
| # The editor materializes default yes/no meanings into an explicit map. | |
| assert all("criteria" in question for question in preset["questions"].values()) | |
| _, actual = read_spec(json.dumps(preset)) | |
| expected = { | |
| name: parse_question(raw) for name, raw in cases[case_name]["questions"].items() | |
| } | |
| assert preset["text"] == cases[case_name]["text"] | |
| assert generation_task(actual) == generation_task(expected) | |
| def test_web_text_and_shared_instructions_reach_one_model_call(playground): | |
| client, backend, _ = playground | |
| spec = request_spec() | |
| spec["instructions"] = "Only use explicit evidence.\nRetain this entire instruction." | |
| spec["questions"]["grade"] = { | |
| "type": "score", | |
| "instructions": "Grade the scene", | |
| "criteria": ["None", "Some", "Many"], | |
| } | |
| spec["questions"]["present"] = { | |
| "type": "independent", | |
| "instructions": "Check presence", | |
| "criteria": {"cat": "A cat", "dog": "A dog"}, | |
| } | |
| response = client.post("/api/run", data={"spec": json.dumps(spec)}) | |
| assert response.status_code == 200 | |
| result = response.json() | |
| assert result["answers"]["animal"]["choice"] == "cat" | |
| assert result["answers"]["grade"]["type"] == "score" | |
| assert set(result["answers"]["present"]["probabilities"]) == {"cat", "dog"} | |
| assert len(backend.calls) == 1 | |
| assert len(backend.calls[0][1]) == 4 | |
| assert all(spec["instructions"] in request.instructions for request in backend.calls[0][1]) | |
| def test_uploaded_image_is_available_during_run_then_removed(playground): | |
| client, backend, directory = playground | |
| image = io.BytesIO() | |
| Image.new("RGB", (8, 8), "red").save(image, format="PNG") | |
| spec = request_spec() | |
| spec["media"] = [{"kind": "image", "name": "../../outside.png"}] | |
| response = client.post( | |
| "/api/run", | |
| data={"spec": json.dumps(spec)}, | |
| files={"media": ("../../outside.png", image.getvalue(), "image/png")}, | |
| ) | |
| assert response.status_code == 200 | |
| assert response.json()["media"][0]["width"] == 8 | |
| assert "attachment-0.png" in backend.calls[0][0].images[0] | |
| assert list(directory.iterdir()) == [] | |
| def test_multi_picture_phone_jpeg_uses_full_resolution_primary_photo(playground): | |
| client, backend, directory = playground | |
| image = io.BytesIO() | |
| Image.new("RGB", (32, 24), "red").save( | |
| image, format="MPO", save_all=True, append_images=[Image.new("RGB", (16, 12), "gray")] | |
| ) | |
| spec = request_spec() | |
| spec["media"] = [{"kind": "image", "name": "phone.jpeg"}] | |
| response = client.post( | |
| "/api/run", | |
| data={"spec": json.dumps(spec)}, | |
| files={"media": ("phone.jpeg", image.getvalue(), "image/jpeg")}, | |
| ) | |
| assert response.status_code == 200 | |
| metadata = response.json()["media"][0] | |
| assert metadata["embedded_images"] == 2 | |
| assert metadata["used_frame"] == 0 | |
| assert (metadata["width"], metadata["height"]) == (32, 24) | |
| assert backend.calls[0][0].images[0].endswith("-primary.png") | |
| assert list(directory.iterdir()) == [] | |
| def test_choice_names_need_no_descriptions_and_grades_need_no_rubric(playground): | |
| client, backend, _ = playground | |
| spec = { | |
| "text": "A Porsche car", | |
| "questions": { | |
| "brand": { | |
| "type": "choice", | |
| "instructions": "Which brand is this?", | |
| "criteria": ["Porsche", "Mercedes"], | |
| }, | |
| "quality": {"type": "score", "instructions": "How well maintained is it?", "levels": 5}, | |
| "tags": { | |
| "type": "independent", | |
| "instructions": "Which are present?", | |
| "criteria": {"car": "", "person": " "}, | |
| }, | |
| }, | |
| } | |
| response = client.post("/api/run", data={"spec": json.dumps(spec)}) | |
| assert response.status_code == 200 | |
| requests = backend.calls[0][1] | |
| assert requests[0].criteria == (("Porsche", "Porsche"), ("Mercedes", "Mercedes")) | |
| assert len(requests[1].criteria) == 5 | |
| assert "lowest" in requests[1].criteria[0][1] | |
| assert response.json()["answers"]["quality"]["type"] == "score" | |
| def test_malformed_media_is_rejected_and_cleaned(playground): | |
| client, backend, directory = playground | |
| spec = request_spec() | |
| spec["media"] = [{"kind": "image", "name": "bad.png"}] | |
| response = client.post( | |
| "/api/run", | |
| data={"spec": json.dumps(spec)}, | |
| files={"media": ("bad.png", b"not an image", "image/png")}, | |
| ) | |
| assert response.status_code == 400 | |
| assert "Could not read this image" in response.json()["error"] | |
| assert backend.calls == [] | |
| assert list(directory.iterdir()) == [] | |
| def test_browser_cannot_submit_local_paths(playground): | |
| client, backend, _ = playground | |
| spec = request_spec() | |
| spec["state"] = {"images": ["/private/file.png"]} | |
| response = client.post("/api/run", data={"spec": json.dumps(spec)}) | |
| assert response.status_code == 400 | |
| assert backend.calls == [] | |
| def test_cross_origin_requests_are_rejected(playground): | |
| client, backend, _ = playground | |
| response = client.post( | |
| "/api/run", | |
| data={"spec": json.dumps(request_spec())}, | |
| headers={"Origin": "https://example.com"}, | |
| ) | |
| assert response.status_code == 403 | |
| assert backend.calls == [] | |
| def test_schema_editor_rejects_duplicates_and_accepts_ordered_grades(playground): | |
| client, _, _ = playground | |
| response = client.post("/api/validate", content='{"questions":{"a":{},"a":{}}}') | |
| assert response.status_code == 400 | |
| assert "Duplicate JSON key" in response.json()["error"] | |
| schema = { | |
| "questions": { | |
| "grade": { | |
| "type": "score", | |
| "instructions": "Grade this", | |
| "criteria": ["Low\nwith details", "High"], | |
| } | |
| } | |
| } | |
| response = client.post("/api/validate", json=schema) | |
| assert response.status_code == 200 | |
| assert response.json() == schema | |
| def test_empty_state_is_rejected_without_inference(playground): | |
| client, backend, _ = playground | |
| spec = request_spec() | |
| spec["text"] = "" | |
| response = client.post("/api/run", data={"spec": json.dumps(spec)}) | |
| assert response.status_code == 400 | |
| assert backend.calls == [] | |
| def test_independent_expansion_is_bounded(): | |
| spec = request_spec() | |
| spec["questions"] = { | |
| "labels": { | |
| "type": "independent", | |
| "instructions": "Check each", | |
| "criteria": {str(i): "An option" for i in range(129)}, | |
| } | |
| } | |
| with pytest.raises(ValueError, match="128 individual fields"): | |
| read_spec(json.dumps(spec)) | |
| def test_static_application_is_served(playground): | |
| client, _, _ = playground | |
| assert "Gemma E2B RLCD" in client.get("/").text | |
| response = client.get("/static/app.js") | |
| assert response.status_code == 200 | |
| assert response.headers["x-content-type-options"] == "nosniff" | |
| assert "Record speech" in response.text | |
| def wav_bytes(seconds): | |
| data = io.BytesIO() | |
| with wave.open(data, "wb") as audio: | |
| audio.setnchannels(1) | |
| audio.setsampwidth(2) | |
| audio.setframerate(16000) | |
| audio.writeframes(b"\0\0" * int(16000 * seconds)) | |
| return data.getvalue() | |
| def test_durationless_webm_recording_is_decoded_in_full(playground): | |
| client, backend, directory = playground | |
| recording = subprocess.run( | |
| [ | |
| "ffmpeg", | |
| "-v", | |
| "error", | |
| "-f", | |
| "wav", | |
| "-i", | |
| "pipe:0", | |
| "-c:a", | |
| "libopus", | |
| "-f", | |
| "webm", | |
| "pipe:1", | |
| ], | |
| input=wav_bytes(0.2), | |
| capture_output=True, | |
| check=True, | |
| ).stdout | |
| spec = request_spec() | |
| spec["text"] = "" | |
| spec["media"] = [{"kind": "video", "name": "recording.webm"}] | |
| response = client.post( | |
| "/api/run", | |
| data={"spec": json.dumps(spec)}, | |
| files={"media": ("recording.webm", recording, "video/webm")}, | |
| ) | |
| assert response.status_code == 200 | |
| assert response.json()["media"][0]["kind"] == "audio" | |
| assert response.json()["media"][0]["duration_seconds"] == pytest.approx(0.2, abs=0.03) | |
| assert backend.calls[0][0].audio[0].endswith("-audio.wav") | |
| assert not backend.calls[0][0].videos | |
| assert list(directory.iterdir()) == [] | |
| def test_long_audio_is_rejected_without_truncation(playground): | |
| client, backend, directory = playground | |
| spec = request_spec() | |
| spec["media"] = [{"kind": "audio", "name": "long.wav"}] | |
| response = client.post( | |
| "/api/run", | |
| data={"spec": json.dumps(spec)}, | |
| files={"media": ("long.wav", wav_bytes(31), "audio/wav")}, | |
| ) | |
| assert response.status_code == 400 | |
| assert "30 seconds" in response.json()["error"] | |
| assert backend.calls == [] | |
| assert list(directory.iterdir()) == [] | |
| def test_concurrent_requests_do_not_overlap_on_model(playground): | |
| client, backend, _ = playground | |
| started, release = Event(), Event() | |
| original = backend.score_batch | |
| def slow_score(state, requests): | |
| started.set() | |
| assert release.wait(5) | |
| return original(state, requests) | |
| backend.score_batch = slow_score | |
| with ThreadPoolExecutor(max_workers=1) as pool: | |
| first = pool.submit(client.post, "/api/run", data={"spec": json.dumps(request_spec())}) | |
| try: | |
| assert started.wait(5) | |
| assert client.get("/api/status").json()["busy"] | |
| second = client.post("/api/run", data={"spec": json.dumps(request_spec())}) | |
| assert second.status_code == 429 | |
| finally: | |
| release.set() | |
| assert first.result(timeout=5).status_code == 200 | |
| assert len(backend.calls) == 1 | |
| def test_comparison_runs_each_path_once_on_same_uploaded_input(playground, monkeypatch): | |
| from pathlib import Path | |
| from gemma_rlcd import comparison | |
| client, backend, directory = playground | |
| normal_calls = [] | |
| def generate(used_backend, state, questions): | |
| assert used_backend is backend | |
| assert Path(state.images[0]).is_file() | |
| normal_calls.append((state, questions)) | |
| return { | |
| "answers": {"animal": "dog"}, | |
| "raw_text": '{"animal":"dog"}', | |
| "valid": True, | |
| "error": None, | |
| "inference_seconds": 0.2, | |
| "output_tokens": 8, | |
| } | |
| monkeypatch.setattr(comparison, "generate_answers", generate) | |
| spec = request_spec() | |
| spec["media"] = [{"kind": "image", "name": "photo.png"}] | |
| image = io.BytesIO() | |
| Image.new("RGB", (8, 8), "red").save(image, format="PNG") | |
| response = client.post( | |
| "/api/compare", | |
| data={"spec": json.dumps(spec)}, | |
| files={"media": ("photo.png", image.getvalue(), "image/png")}, | |
| ) | |
| assert response.status_code == 200 | |
| data = response.json() | |
| details = data["comparison"] | |
| assert len(backend.calls) == len(normal_calls) == 1 | |
| assert backend.calls[0][0] == normal_calls[0][0] | |
| assert details["agreement"] == {"animal": False} | |
| assert details["parallel_values"] == {"animal": "cat"} | |
| assert details["seconds"]["normal"] == pytest.approx(0.2 + data["media_prepare_seconds"]) | |
| assert details["seconds"]["parallel"] == pytest.approx( | |
| data["decision_seconds"] + data["media_prepare_seconds"] | |
| ) | |
| assert details["normal_over_parallel"] > 0 | |
| assert data["answers"]["animal"]["choice"] == "cat" | |
| assert list(directory.iterdir()) == [] | |
| def test_invalid_generated_answer_remains_visible_without_speedup_claim(playground, monkeypatch): | |
| from gemma_rlcd import comparison | |
| client, _, _ = playground | |
| monkeypatch.setattr( | |
| comparison, | |
| "generate_answers", | |
| lambda *args: { | |
| "answers": None, | |
| "raw_text": "not json", | |
| "valid": False, | |
| "error": "Invalid JSON", | |
| "inference_seconds": 0.1, | |
| "output_tokens": 2, | |
| }, | |
| ) | |
| response = client.post("/api/compare", data={"spec": json.dumps(request_spec())}) | |
| assert response.status_code == 200 | |
| data = response.json() | |
| assert data["comparison"]["normal_over_parallel"] is None | |
| assert data["comparison"]["agreement"]["animal"] is None | |
| assert data["comparison"]["normal"]["raw_text"] == "not json" | |
| assert data["answers"]["animal"]["choice"] == "cat" | |
| def test_visual_demo_has_128_distinct_checks_and_supports_each_size(playground): | |
| client, backend, _ = playground | |
| assert "One scene." in client.get("/demo").text | |
| config = client.get("/static/visual-demo.json").json() | |
| assert len(config["groups"]) == 4 | |
| for size in (32, 64, 128): | |
| spec = { | |
| "text": "scene", | |
| "instructions": config["instructions"], | |
| "questions": { | |
| group["id"]: { | |
| "type": "independent", | |
| "instructions": "Which checks are visible?", | |
| "criteria": { | |
| check["id"]: check["description"] for check in group["checks"][: size // 4] | |
| }, | |
| } | |
| for group in config["groups"] | |
| }, | |
| } | |
| response = client.post("/api/run", data={"spec": json.dumps(spec)}) | |
| assert response.status_code == 200, response.text | |
| assert len(backend.calls[-1][1]) == size | |
| def test_streaming_comparison_orders_events_and_cleans_uploads(playground, monkeypatch): | |
| from gemma_rlcd import comparison | |
| client, _, directory = playground | |
| def generate(backend, state, questions, on_token=None, on_progress=None): | |
| on_token('{"animal":', 1) | |
| on_token('"cat"}', 2) | |
| return { | |
| "answers": {"animal": "cat"}, | |
| "raw_text": '{"animal":"cat"}', | |
| "valid": True, | |
| "error": None, | |
| "inference_seconds": 1, | |
| "output_tokens": 2, | |
| } | |
| monkeypatch.setattr(comparison, "generate_answers", generate) | |
| response = client.post("/api/compare-stream", data={"spec": json.dumps(request_spec())}) | |
| assert response.status_code == 200 | |
| assert response.headers["content-type"].startswith("application/x-ndjson") | |
| events = [json.loads(line) for line in response.text.splitlines()] | |
| assert [event["type"] for event in events[:2]] == ["accepted", "race_start"] | |
| assert events[-1]["type"] == "complete" | |
| assert [event["type"] for event in events if event.get("method") == "parallel"] == [ | |
| "phase_start", | |
| "answer", | |
| "phase_complete", | |
| ] | |
| assert [event["type"] for event in events if event.get("method") == "normal"] == [ | |
| "phase_start", | |
| "token", | |
| "token", | |
| "phase_complete", | |
| ] | |
| assert next(event for event in events if event["type"] == "answer")["value"] == "cat" | |
| assert events[-1]["result"]["comparison"]["normal"]["valid"] | |
| assert events[-1]["result"]["comparison"]["methodology"]["execution"] == "concurrent_shared_gpu" | |
| assert not client.get("/api/status").json()["busy"] | |
| assert list(directory.iterdir()) == [] | |
| def test_stream_validation_failure_releases_gpu_lock(playground): | |
| client, _, directory = playground | |
| response = client.post("/api/compare-stream", data={"spec": "{}"}) | |
| assert response.status_code == 400 | |
| assert not client.get("/api/status").json()["busy"] | |
| assert list(directory.iterdir()) == [] | |