botp
/

File size: 6,658 Bytes
1d2de8a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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