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_json_backend.py from larkooo/gemma-e2b-rlcd: direct link, hf CLI and curl.
- Browser
- Download file 3.08 kB
-
https://huggingface.co/larkooo/gemma-e2b-rlcd/resolve/main/tests/test_json_backend.py
- Command line
-
hf download hf://larkooo/gemma-e2b-rlcd/tests/test_json_backend.py
-
curl -L -o test_json_backend.py https://huggingface.co/larkooo/gemma-e2b-rlcd/resolve/main/tests/test_json_backend.py
3.08 kB
| import math | |
| from types import SimpleNamespace | |
| import pytest | |
| from gemma_rlcd.json_backend import JSONMLXBackend | |
| from gemma_rlcd.json_scoring import FieldTokens | |
| mx = pytest.importorskip("mlx.core") | |
| def test_complete_candidate_likelihood_uses_suffix_and_ignores_padding(batch_size): | |
| backend = JSONMLXBackend.__new__(JSONMLXBackend) | |
| backend.mx = mx | |
| backend.branch_batch_size = batch_size | |
| backend.tokenizer = SimpleNamespace(pad_token_id=0) | |
| backend.fork_cache = lambda cache, size: [] | |
| table = [[-30.0] * 16 for _ in range(16)] | |
| table[9][1], table[9][4] = math.log(0.9), math.log(0.1) | |
| table[1][2], table[1][3] = math.log(0.1), math.log(0.9) | |
| transition = mx.array(table) | |
| backend.model = SimpleNamespace( | |
| language_model=SimpleNamespace( | |
| model=lambda *, inputs, cache, logits_to_keep: inputs[:, -logits_to_keep:, None], | |
| logits_from_hidden=lambda hidden: transition[hidden[..., 0]], | |
| ) | |
| ) | |
| fields = [(0, FieldTokens((9,), ((1, 2), (1, 3), (4,))))] | |
| streamed = [] | |
| results, batches = backend._sequence_scores([], 100, fields, on_scores=streamed.append) | |
| likelihoods = [math.exp(value) for value in results[0].logits] | |
| assert likelihoods == pytest.approx([0.09, 0.81, 0.1], abs=1e-6) | |
| assert results[0].allowed_token_mass == pytest.approx(1, abs=1e-6) | |
| assert sum(batches) == 3 | |
| assert max(batches) <= batch_size | |
| assert streamed == [[(0, results[0])]] | |
| def test_repeated_schema_reuses_only_tokenization_and_still_prefills_each_image( | |
| monkeypatch, tmp_path | |
| ): | |
| from gemma_rlcd import json_backend | |
| from gemma_rlcd.core import Noul, State, TokenScores | |
| encoded, prefills = [], [] | |
| backend = JSONMLXBackend.__new__(JSONMLXBackend) | |
| backend.tokenizer = SimpleNamespace( | |
| encode=lambda text, **kwargs: encoded.append(text) or list(map(ord, text)) | |
| ) | |
| backend.max_input_tokens = 8192 | |
| backend.last_stats = {} | |
| monkeypatch.setattr( | |
| json_backend, | |
| "prepare_generation", | |
| lambda backend, state, questions: ( | |
| "prompt", | |
| {"input_ids": SimpleNamespace(shape=(1, 6)), "image": state.images}, | |
| ), | |
| ) | |
| backend.prefill = lambda prepared: prefills.append(prepared.inputs["image"]) or [] | |
| backend._sequence_scores = lambda cache, tokens, fields, on_scores: ( | |
| {index: TokenScores((float(len(prefills)), 0), 1, 10) for index, _ in fields}, | |
| [2], | |
| ) | |
| questions = {"visible": Noul("Visible?")} | |
| paths = [tmp_path / name for name in ("one.jpg", "two.jpg")] | |
| for path in paths: | |
| path.touch() | |
| first = backend.score_questions(State(images=(str(paths[0]),)), questions) | |
| assert not backend.last_stats["schema_cache_hit"] | |
| count = len(encoded) | |
| second = backend.score_questions(State(images=(str(paths[1]),)), questions) | |
| assert backend.last_stats["schema_cache_hit"] | |
| assert len(encoded) == count | |
| assert prefills == [(str(paths[0]),), (str(paths[1]),)] | |
| assert first[0].logits != second[0].logits | |