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
File size: 3,649 Bytes
53e24ca | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 | """Paired full-detail execution comparison against version 0.3 computations."""
import argparse
import json
import statistics
import time
from pathlib import Path
from benchmark_head import requests_for
from check_cache import compare
from mlx_vlm.models.gemma4.vision import VisionPatchEmbedder
from gemma_rlcd import State
from gemma_rlcd.cached_backend import CachedMLXBackend
from gemma_rlcd.core import softmax
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--model", required=True)
parser.add_argument("--media", required=True, type=Path)
parser.add_argument("--report", required=True, type=Path)
parser.add_argument("--repeats", type=int, default=5)
args = parser.parse_args()
backend = CachedMLXBackend(args.model, branch_batch_size=32)
indexed = backend.model.vision_tower.patch_embedder
dense = VisionPatchEmbedder(backend.model.config.vision_config)
dense.input_proj = indexed.input_proj
dense.position_embedding_table = indexed.position_embedding_table
media = args.media.resolve()
jobs = [
("text", State(text="A cat sleeps on a sofa. No dogs are present."), 4),
("text_many_fields", State(text="A cat sleeps on a sofa. No dogs are present."), 28),
("image", State(images=(str(media / "red.png"),)), 4),
("speech", State(audio=(str(media / "dog.wav"),)), 4),
("video", State(videos=(str(media / "red-blue.mp4"),)), 4),
("video_speech", State(videos=(str(media / "red-blue-speech.mp4"),)), 4),
]
report = {
"status": "paired_full_detail_execution_probe_not_quality_benchmark",
"compute_dtype": "float32",
"state_layers": 35,
"timing_scope": "warm_model_fresh_state_including_preprocessing_and_all_fields",
"reference": "v0.3_native_answer_tail_and_dense_vision_positions",
"optimized": "per_row_answer_gather_and_indexed_vision_positions",
"results": [],
}
for name, state, count in jobs:
requests = requests_for(backend, count)
def run(mode):
backend.answer_mode = "tail" if mode == "reference" else "gather"
backend.model.vision_tower.patch_embedder = dense if mode == "reference" else indexed
start = time.perf_counter()
scores = backend.score_batch(state, requests)
for score in scores:
softmax(score.logits)
elapsed = time.perf_counter() - start
return scores, elapsed, dict(backend.last_stats)
modes = ["reference", "optimized"]
for mode in modes:
run(mode)
records = {mode: [] for mode in modes}
scores = {}
stats = {}
for repeat in range(args.repeats):
for mode in modes if repeat % 2 == 0 else reversed(modes):
scores[mode], elapsed, stats[mode] = run(mode)
records[mode].append(elapsed)
medians = {mode: statistics.median(values) for mode, values in records.items()}
result = {
"state": name,
"fields": count,
"comparison": compare(scores["reference"], scores["optimized"]),
"raw_seconds": records,
"median_seconds": medians,
"speed_ratio": medians["reference"] / medians["optimized"],
"execution": stats,
}
report["results"].append(result)
print(
json.dumps({key: value for key, value in result.items() if key != "execution"}),
flush=True,
)
args.report.write_text(json.dumps(report, indent=2) + "\n")
if __name__ == "__main__":
main()
|