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] hf download larkooo/gemma-e2b-rlcd --local-dir gemma-e2b-rlcd
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
Download tests/test_comparison.py from larkooo/gemma-e2b-rlcd: direct link, hf CLI and curl.
- Browser
- Download file 6.34 kB
-
https://huggingface.co/larkooo/gemma-e2b-rlcd/resolve/main/tests/test_comparison.py
- Command line
-
hf download hf://larkooo/gemma-e2b-rlcd/tests/test_comparison.py
-
curl -L -o test_comparison.py https://huggingface.co/larkooo/gemma-e2b-rlcd/resolve/main/tests/test_comparison.py
6.34 kB
| import json | |
| import pytest | |
| from gemma_rlcd.comparison import discrete_answers, generation_task, parse_generated | |
| from gemma_rlcd.core import Choice, Independent, Noul, Score, State | |
| def questions(): | |
| return { | |
| "brand": Choice("Which brand is this?", {"Porsche": "Porsche", "Mercedes": "Mercedes"}), | |
| "quality": Score("Grade its condition", ["Poor", "Good", "Excellent"]), | |
| "moving": Noul("Is it moving?"), | |
| "presence": Independent("Which are present?", {"car": "A car", "person": "A person"}), | |
| } | |
| def test_generation_requests_all_fields_in_one_json_answer(): | |
| task = generation_task(questions()) | |
| assert "Porsche" in task and "Mercedes" in task | |
| assert "Grade its condition" in task and "Excellent" in task | |
| assert "a boolean" in task and "proposition name to a boolean" in task | |
| expected = { | |
| "brand": "Porsche", | |
| "quality": 2, | |
| "moving": False, | |
| "presence": {"car": True, "person": False}, | |
| } | |
| assert parse_generated(json.dumps(expected), questions()) == expected | |
| assert parse_generated("```json\n" + json.dumps(expected) + "\n```", questions()) == expected | |
| def test_generated_output_must_match_the_contract(text): | |
| with pytest.raises(ValueError): | |
| parse_generated(text, questions()) | |
| def test_agreement_uses_grade_mode_and_boolean_thresholds(): | |
| assert discrete_answers( | |
| { | |
| "grade": { | |
| "type": "score", | |
| "score": 1.4, | |
| "probabilities": {"0": 0.1, "1": 0.4, "2": 0.5}, | |
| }, | |
| "moving": {"type": "noul", "noul": 0.7}, | |
| "presence": {"type": "independent", "probabilities": {"car": 0.9, "person": 0.1}}, | |
| } | |
| ) == {"grade": 2, "moving": True, "presence": {"car": True, "person": False}} | |
| def test_streaming_generation_uses_real_chunks_and_validates_complete_answer(monkeypatch): | |
| import sys | |
| from types import SimpleNamespace | |
| from gemma_rlcd import comparison | |
| seen = [] | |
| parts = ['{"visible":', "true", "}"] | |
| def stream(*args, **kwargs): | |
| for index, text in enumerate(parts): | |
| assert len(seen) == index | |
| yield SimpleNamespace( | |
| text=text, | |
| generation_tokens=index + 1, | |
| prompt_tokens=50, | |
| finish_reason="stop" if index == len(parts) - 1 else None, | |
| ) | |
| monkeypatch.setitem( | |
| sys.modules, "mlx_vlm", SimpleNamespace(generate=None, stream_generate=stream) | |
| ) | |
| monkeypatch.setattr(comparison, "prepare_generation", lambda *args: ("prompt", {})) | |
| backend = SimpleNamespace( | |
| model=None, | |
| processor=None, | |
| tokenizer=SimpleNamespace(encode=lambda *args, **kwargs: [1]), | |
| mx=SimpleNamespace(synchronize=lambda: None), | |
| ) | |
| result = comparison.generate_answers( | |
| backend, | |
| State(text="test"), | |
| {"visible": Noul("Visible?")}, | |
| on_token=lambda text, count: seen.append((text, count)), | |
| ) | |
| assert seen == list(zip(parts, [1, 2, 3], strict=True)) | |
| assert result["valid"] and result["answers"] == {"visible": True} | |
| assert result["raw_text"] == "".join(parts) | |
| def test_concurrent_comparison_streams_both_paths_before_either_finishes(monkeypatch): | |
| from threading import Event | |
| from types import SimpleNamespace | |
| from gemma_rlcd import comparison | |
| from gemma_rlcd.core import TokenScores | |
| scored, generated = Event(), Event() | |
| events = [] | |
| class Backend: | |
| last_stats = {} | |
| processor = SimpleNamespace(tokenizer=SimpleNamespace(mutable=[])) | |
| def symbols(self, count): | |
| return ("A", "B") | |
| def score_questions(self, state, questions, on_scores, on_progress=None): | |
| assert generated.wait(2), "Generation never started alongside scoring" | |
| result = TokenScores((5, 0), 1, 20) | |
| on_scores([(0, result)]) | |
| scored.set() | |
| return [result] | |
| backend = Backend() | |
| def generate(worker, state, questions, on_token, on_progress=None): | |
| assert worker.processor is not backend.processor | |
| assert worker.processor.tokenizer is not backend.processor.tokenizer | |
| on_token('{"visible":', 1) | |
| generated.set() | |
| assert scored.wait(2), "Scoring never completed while generation was active" | |
| on_token("true}", 2) | |
| return {"answers": {"visible": True}, "valid": True, "inference_seconds": 0} | |
| monkeypatch.setattr(comparison, "generate_answers", generate) | |
| _, result = comparison.compare( | |
| backend, | |
| State(text="scene"), | |
| {"visible": Noul("Visible?")}, | |
| 0, | |
| emit=events.append, | |
| concurrent=True, | |
| ) | |
| assert events[0]["type"] == "race_start" | |
| first_token = next(i for i, event in enumerate(events) if event["type"] == "token") | |
| first_answer = next(i for i, event in enumerate(events) if event["type"] == "answer") | |
| first_finish = next(i for i, event in enumerate(events) if event["type"] == "phase_complete") | |
| assert first_token < first_answer < first_finish | |
| assert result["agreement"] == {"visible": True} | |
| assert result["methodology"]["execution"] == "concurrent_shared_gpu" | |
| def test_concurrent_comparison_can_stop_before_or_after_worker_launch(cancel_at): | |
| from types import SimpleNamespace | |
| from gemma_rlcd.comparison import compare | |
| def emit(event): | |
| if event["type"] == cancel_at: | |
| raise RuntimeError("Client stopped") | |
| with pytest.raises(RuntimeError, match="Client stopped"): | |
| compare( | |
| SimpleNamespace(), | |
| State(text="scene"), | |
| {"visible": Noul("Visible?")}, | |
| 0, | |
| emit=emit, | |
| concurrent=True, | |
| ) | |