Text Classification
PEFT
lora
document-question-answering
structured-decisions
calibration
synthetic-evaluation
Instructions to use DoccyHealth/Solomon with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use DoccyHealth/Solomon with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
File size: 6,658 Bytes
5c0a4a8 | 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 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 | """End-to-end BF16 library wiring on a synthetic hybrid vision model.
This is an integration test, not evidence about Solomon's trained accuracy.
Only checkpoint hashes and the tokenizer are replaced in the test fixture.
"""
import json
from dataclasses import asdict
from types import SimpleNamespace
from typing import ClassVar
import mlx.core as mx
import numpy as np
from mlx.utils import tree_map_with_path
from mlx_vlm.models.qwen3_5.config import ModelConfig, TextConfig, VisionConfig
from mlx_vlm.models.qwen3_5.qwen3_5 import Model
from mlx_vlm.models.qwen3_vl.processing_qwen3_vl import Qwen3VLImageProcessor, Qwen3VLProcessor
from mlx_vlm.utils import save_weights
from PIL import Image
import solomon_mlx.engine as engine_module
from solomon_mlx import Solomon
from solomon_mlx.artifacts import BASE_REVISION, SOLOMON_REVISION, sha256
from solomon_mlx.prepare import convert_bf16
class Tokenizer:
markers: ClassVar[dict] = {"<|vision_start|>": 251, "<|vision_end|>": 252, "<|image_pad|>": 253}
def apply_chat_template(self, messages, **kwargs):
return "\n".join(m["content"] for m in messages) + "\nAssistant:"
def encode(self, text, **kwargs):
for key, value in self.markers.items():
text = text.replace(key, chr(value))
return list(text.encode("latin1"))
def convert_tokens_to_ids(self, text):
return self.markers[text]
def test_full_library_bf16_text_image_repeated_question(tmp_path, monkeypatch):
mx.random.seed(14)
text = TextConfig(
model_type="qwen3_5_text",
hidden_size=5120,
intermediate_size=64,
linear_num_value_heads=2,
linear_num_key_heads=2,
linear_key_head_dim=32,
linear_value_head_dim=32,
linear_conv_kernel_dim=4,
num_hidden_layers=4,
num_attention_heads=4,
num_key_value_heads=2,
head_dim=32,
rms_norm_eps=1e-6,
vocab_size=256,
max_position_embeddings=4096,
rope_parameters={
"type": "default",
"mrope_section": [2, 1, 1],
"rope_theta": 100000,
"partial_rotary_factor": 0.25,
},
)
vision = VisionConfig(
depth=1,
hidden_size=64,
intermediate_size=64,
out_hidden_size=5120,
num_heads=4,
patch_size=16,
spatial_patch_size=16,
num_position_embeddings=64,
deepstack_visual_indexes=[],
)
config = ModelConfig(
text_config=text,
vision_config=vision,
model_type="qwen3_5",
image_token_id=253,
video_token_id=254,
vision_start_token_id=251,
vision_end_token_id=252,
vocab_size=256,
)
model = Model(config)
model.update(
tree_map_with_path(
lambda k, v: v.astype(mx.float32 if v.ndim == 1 else mx.bfloat16), model.parameters()
)
)
backbone = tmp_path / "backbone"
save_weights(backbone, model, donate_weights=True)
(backbone / "config.json").write_text(json.dumps(asdict(config)))
source = tmp_path / "original"
backbone.rename(source)
convert_bf16(source, backbone)
# Cloud conversion uses the CPU backend. Its output must match the Mac's
# default backend before either is used by the same Metal runtime.
cpu_backbone = tmp_path / "cpu-backbone"
with mx.stream(mx.cpu):
convert_bf16(source, cpu_backbone)
for shard in backbone.glob("*.safetensors"):
gpu_arrays = mx.load(str(shard))
cpu_arrays = mx.load(str(cpu_backbone / shard.name))
assert gpu_arrays.keys() == cpu_arrays.keys()
for key in gpu_arrays:
assert gpu_arrays[key].dtype == cpu_arrays[key].dtype
assert mx.array_equal(gpu_arrays[key], cpu_arrays[key]).item(), key
adapter = tmp_path / "adapter.safetensors"
mx.save_safetensors(
str(adapter),
{
"model.layers.0.mlp.gate_proj.lora_a": mx.full((5120, 64), 0.001, mx.float32),
"model.layers.0.mlp.gate_proj.lora_b": mx.full((64, 64), 0.001, mx.float32),
},
)
keys = [
"boolean/state4",
"entity/state4",
"multilabel/state4",
"ordered/threshold4",
"single/choiceR",
"single/choiceS",
"single/sufficiency3",
"ordered/choiceR",
"ordered/choiceS",
"ordered/sufficiency3",
]
rng = np.random.default_rng(14)
heads = {}
for k in keys:
heads[k + "/weight"] = rng.normal(0, 0.01, (10, 5120)).astype(np.float32)
heads[k + "/bias"] = np.zeros(10, np.float32)
np.savez(tmp_path / "heads.npz", **heads)
binding = {
"profile": "quality",
"schema": "solomon-mlx-binding-v1",
"base_revision": BASE_REVISION,
"solomon_revision": SOLOMON_REVISION,
"dtype": "bfloat16",
"files": {str(p.relative_to(tmp_path)): sha256(p) for p in tmp_path.rglob("*") if p.is_file()},
}
(tmp_path / "binding.json").write_text(json.dumps(binding))
monkeypatch.setattr(engine_module, "ADAPTER_SHA", sha256(adapter))
monkeypatch.setattr(engine_module, "HEADS_SHA", sha256(tmp_path / "heads.npz"))
processor = SimpleNamespace(
tokenizer=Tokenizer(), image_processor=Qwen3VLImageProcessor(min_pixels=1024, max_pixels=16384)
)
monkeypatch.setattr(Qwen3VLProcessor, "from_pretrained", lambda *a, **kw: processor)
port = Solomon.load(tmp_path, chunk_size=128)
page = tmp_path / "page.png"
Image.new("RGB", (128, 128), "white").save(page)
questions = {
"a": "Is Alice certified?",
"b": {"type": "choice", "instructions": "Who?", "options": ["Alice", "Bob"]},
"c": {"type": "score", "instructions": "Level?", "levels": ["low", "high"]},
"d": {"instructions": "Is {candidate} certified?", "candidates": ["Alice", "Bob"]},
"e": {"instructions": "Which labels apply?", "candidates": ["certified", "unavailable"]},
}
with port.prefill([{"text": "Alice is certified."}, {"image": str(page)}, {"text": "End."}]) as state:
assert state._data["counts"] == [16]
first = port.decide(state=state, questions=questions, evidence="none", diagnostics=True)
repeated = port.decide(
state=state, questions={"a": questions["a"]}, evidence="support", diagnostics=True
)
assert first["answers"]["a"]["noul"] == repeated["answers"]["a"]["noul"]
assert repeated["answers"]["a"]["evidence_status"] == "unsupported_page_selector"
assert set(first["answers"]) == set(questions)
assert port.engine.context["start"] is None
|