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: 4,914 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 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 | """Full warm-request runtime probes. Untrained head outputs are not quality evidence."""
import argparse
import json
import statistics
from pathlib import Path
from gemma_rlcd import State
from gemma_rlcd.core import ScoringRequest, decision_prompt
from gemma_rlcd.head_backend import DecisionHeadBackend
def requests_for(backend, count):
tasks = [
(
"Which animal is mentioned or visible?",
{"cat": "A cat", "dog": "A dog", "other": "Neither"},
),
("Is a dog mentioned?", {"true": "Yes", "false": "No"}),
("How much red is visible?", {"0": "No red", "1": "Some red", "2": "Mostly red"}),
("Is a sofa mentioned?", {"true": "Yes", "false": "No"}),
]
requests = []
for i in range(count):
instruction, criteria = tasks[i % len(tasks)]
symbols = tuple(backend.symbols(len(criteria)))
requests.append(
ScoringRequest(
decision_prompt(instruction, criteria, symbols),
symbols,
instruction,
tuple(criteria.items()),
)
)
return requests
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("--state-layers", type=int, default=35)
parser.add_argument("--image-soft-tokens", type=int, default=280)
parser.add_argument("--video-max-frames", type=int, default=32)
parser.add_argument("--dtype", choices=["float16", "float32", "bfloat16"], default="float32")
parser.add_argument("--fields", type=int, nargs="+", default=[1, 4, 16, 28])
parser.add_argument("--repeats", type=int, default=5)
parser.add_argument(
"--states", nargs="+", default=["text", "image", "speech", "video", "video_speech"]
)
args = parser.parse_args()
if args.repeats < 1:
parser.error("repeats must be positive")
media = args.media.resolve()
states = {
"text": State(text="A cat sleeps on a sofa. No dogs are present."),
"image": State(images=(str(media / "red.png"),)),
"speech": State(audio=(str(media / "dog.wav"),)),
"video": State(videos=(str(media / "red-blue.mp4"),)),
"video_speech": State(videos=(str(media / "red-blue-speech.mp4"),)),
}
backend = DecisionHeadBackend(
args.model,
state_layers=args.state_layers,
compute_dtype=args.dtype,
image_soft_tokens=args.image_soft_tokens,
video_max_frames=args.video_max_frames,
)
report = {
"status": "untrained_runtime_probe_not_quality_or_calibration_evidence",
"model": args.model,
"state_layers": args.state_layers,
"compute_dtype": args.dtype,
"head_config": backend.head.config.to_dict(),
"image_soft_tokens": args.image_soft_tokens,
"video_max_frames": args.video_max_frames,
"timing_scope": "warm_model_fresh_state_including_media_decode_preprocess_encoding_and_all_fields",
"field_scaling": "four representative question templates repeated to the requested count",
"results": [],
}
for name in args.states:
state = states[name]
requests = {count: requests_for(backend, count) for count in args.fields}
for count in args.fields:
backend.probe(state, requests[count])
samples = {count: [] for count in args.fields}
for repeat in range(args.repeats):
for count in args.fields if repeat % 2 == 0 else reversed(args.fields):
try:
samples[count].append(backend.probe(state, requests[count]))
except Exception as exc:
samples[count].append({"error": f"{type(exc).__name__}: {exc}"})
for count, records in samples.items():
successful = [record for record in records if "error" not in record]
record = {
"state": name,
"fields": count,
"attempted": len(records),
"completed": len(successful),
"samples": records,
}
if successful:
timing_keys = [key for key in successful[0] if key.endswith("_seconds")]
record["median_seconds"] = {
key: statistics.median(row[key] for row in successful) for key in timing_keys
}
record["max_seconds"] = max(row["total_seconds"] for row in successful)
report["results"].append(record)
print(
json.dumps({key: value for key, value in record.items() if key != "samples"}),
flush=True,
)
args.report.write_text(json.dumps(report, indent=2) + "\n")
if __name__ == "__main__":
main()
|