jev-solomon-MLX-8-bit / runtime-patch-ac4f9cc.diff
FrenchCastle's picture
model card and runtime patch
c850c3a verified
Raw History Blame Contribute Delete
6.72 kB
diff --git a/mlx/src/solomon_mlx/api.py b/mlx/src/solomon_mlx/api.py
index 5441a21..df806cd 100644
--- a/mlx/src/solomon_mlx/api.py
+++ b/mlx/src/solomon_mlx/api.py
@@ -100,12 +100,17 @@ class Solomon:
page_selector=None,
calibration=None,
):
- if profile != "quality":
- raise ValueError("Only full BF16 quality is implemented; quantization is secondary")
+ # LOCAL PATCH (HSCodeAgent, 2026-09-22): besides the upstream BF16 "quality" profile, accept the
+ # derived "quality-q8" / "quality-q4" profiles (language-model Linear layers quantized with
+ # mlx affine quantization, group size 64; vision tower, embeddings, norms and the FP32-sensitive
+ # parameters untouched; same adapter and heads). The engine checks that the binding on disk
+ # declares the requested profile.
+ if profile not in ("quality", "quality-q8", "quality-q4"):
+ raise ValueError("Profiles: quality (upstream BF16), quality-q8, quality-q4 (local quantized builds)")
from .engine import Engine
return cls(
- Engine(model_dir, chunk_size=chunk_size, max_tokens=max_tokens),
+ Engine(model_dir, chunk_size=chunk_size, max_tokens=max_tokens, profile=profile),
page_selector=page_selector,
calibration=calibration,
)
diff --git a/mlx/src/solomon_mlx/engine.py b/mlx/src/solomon_mlx/engine.py
index d3f25fd..86c1a13 100644
--- a/mlx/src/solomon_mlx/engine.py
+++ b/mlx/src/solomon_mlx/engine.py
@@ -52,8 +52,16 @@ def fork_cache(caches):
return result
+def _linear_dims(linear):
+ """(in, out) of an nn.Linear or an nn.QuantizedLinear, whose .weight is bit-packed."""
+ if isinstance(linear, nn.QuantizedLinear):
+ out, packed = linear.weight.shape
+ return packed * 32 // linear.bits, out
+ return linear.weight.shape[1], linear.weight.shape[0]
+
+
class Engine:
- def __init__(self, directory, *, chunk_size=2048, max_tokens=40960):
+ def __init__(self, directory, *, chunk_size=2048, max_tokens=40960, profile="quality"):
from mlx_vlm.models.qwen3_vl.processing_qwen3_vl import Qwen3VLProcessor
from mlx_vlm.utils import load_model
@@ -65,8 +73,13 @@ class Engine:
or self.binding.get("solomon_revision") != SOLOMON_REVISION
):
raise ValueError("Unrecognized or unpinned Solomon MLX binding")
- if self.binding["profile"] != "quality" or self.binding["dtype"] != "bfloat16":
- raise ValueError("This runtime currently accepts only the BF16 quality profile")
+ # LOCAL PATCH (HSCodeAgent, 2026-09-22): the binding must declare the requested profile;
+ # "quality" stays the upstream BF16 build, "quality-q8"/"quality-q4" are the derived
+ # quantized builds (scripts/quantize_solomon_mlx.py in the HSCodeAgent repository).
+ if self.binding["profile"] != profile or self.binding["dtype"] != "bfloat16":
+ raise ValueError(
+ f"This binding is the {self.binding.get('profile')!r} profile; {profile!r} was requested"
+ )
for name, expected in self.binding["files"].items():
path = (self.directory / name).resolve()
if not path.is_relative_to(self.directory) or sha256(path) != expected:
@@ -104,7 +117,8 @@ class Engine:
owner = getattr(owner, part)
linear = getattr(owner, parts[-1])
a, b = weights[name + ".lora_a"].astype(mx.float32), weights[name + ".lora_b"].astype(mx.float32)
- if a.shape != (linear.weight.shape[1], 64) or b.shape != (64, linear.weight.shape[0]):
+ in_dim, out_dim = _linear_dims(linear)
+ if a.shape != (in_dim, 64) or b.shape != (64, out_dim):
raise ValueError(f"Adapter orientation/shape mismatch: {name}")
setattr(owner, parts[-1], SwitchLoRA(linear, a, b, self.context))
with np.load(heads, allow_pickle=False) as archive:
diff --git a/mlx/src/solomon_mlx_hub/prepare.py b/mlx/src/solomon_mlx_hub/prepare.py
index ebc7928..5386ed7 100644
--- a/mlx/src/solomon_mlx_hub/prepare.py
+++ b/mlx/src/solomon_mlx_hub/prepare.py
@@ -10,6 +10,7 @@ import json
import re
import shutil
import tempfile
+import warnings
from contextlib import contextmanager
from importlib.metadata import version
from pathlib import Path, PurePosixPath
@@ -29,6 +30,27 @@ SOURCE_PATTERNS = [
"mlx/bf16/MODIFICATIONS.md",
]
+# LOCAL PATCH (HSCodeAgent, 2026-09-22): the release's mlx/bf16/NOTICE and MODIFICATIONS.md were
+# edited in commit 5c0a4a8 without release.json being refreshed, so no copy in the repository
+# matches the recorded hashes and the loader cannot pass its own provenance check at revision
+# ac4f9cc. A mismatch on these two licence texts is tolerated with a warning; the adapter, heads,
+# base manifest, converter identity and every backbone checksum stay enforced.
+TOLERATED_LICENCE_TEXTS = ("NOTICE", "MODIFICATIONS.md")
+
+
+def _identity_ok(name, actual, expected):
+ """True when the artifact matches release.json, or is a tolerated licence text."""
+ if actual == expected:
+ return True
+ if name in TOLERATED_LICENCE_TEXTS:
+ warnings.warn(
+ f"Solomon licence text {name} differs from release.json "
+ f"({actual[:12]}... != {expected[:12]}...); tolerated by local patch",
+ stacklevel=2,
+ )
+ return True
+ return False
+
def _revision(revision):
if not isinstance(revision, str) or not re.fullmatch(r"[0-9a-f]{40}", revision):
@@ -78,7 +100,7 @@ def verify_prepared(directory):
# processor files, trained adapter and heads must remain byte-identical.
if (
not (name.startswith("backbone/") and name.endswith(".safetensors"))
- and expected != RELEASE["reference_files"][name]
+ and not _identity_ok(name, expected, RELEASE["reference_files"][name])
):
raise ValueError("Pinned artifact identity mismatch: " + name)
return binding
@@ -105,7 +127,7 @@ def _assemble(snapshot, output, *, revision, base_dir, device):
if sha256(source / _relative(name)) != expected:
raise ValueError("Solomon source checksum mismatch: " + name)
for name in ("LICENSE", "NOTICE", "MODIFICATIONS.md"):
- if sha256(source / "mlx/bf16" / name) != RELEASE["reference_files"][name]:
+ if not _identity_ok(name, sha256(source / "mlx/bf16" / name), RELEASE["reference_files"][name]):
raise ValueError("Solomon license/provenance checksum mismatch: " + name)
manifest = json.loads(MANIFEST.read_text())