Text Classification
PEFT
lora
document-question-answering
structured-decisions
calibration
synthetic-evaluation
Instructions to use botp/Solomon with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use botp/Solomon with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
| """The answer contract engine: contract v3 on the reference engine's float32 cached path. | |
| Adds to solomon/engine_reference.py (which is not modified): | |
| * a switchable LoRA layer, so an adapter can act on every token ('full', as the reference engine) or only on the | |
| question branch ('question'): off while the document is prefilled, on for the branch, and in full | |
| execution applied from the prefix boundary onward; | |
| * documents made of parts: text and page images, in any order. The vision tower runs once per document, | |
| the library's positions are used for the prefix, and the state carries the rope delta; | |
| * task blocks for every answer type, and ten answer letters; | |
| * state format v3. | |
| """ | |
| import copy | |
| import hashlib | |
| import json | |
| from pathlib import Path | |
| import numpy as np | |
| from solomon.engine_reference import (BOOLEAN_TASK, CHOICE_TASK, CHUNK, SYSTEM, TOKEN_CAP, Engine, Ledger, boolean_block, # noqa: F401 | |
| choice_block, dump, read_rows, sha, softmax) | |
| ROOT = Path(__file__).resolve().parents[1] | |
| MODEL = ROOT / 'runtime/qwen38-8bit' | |
| ART = ROOT / 'artifacts/contract' | |
| CONTRACT = 'solomon-answer-contract-v3' | |
| STATE_FORMAT = 'solomon-answer-state-v3' | |
| LETTERS = 'ABCDEFGHIJ' | |
| PAGE = '<|vision_start|><|image_pad|><|vision_end|>' | |
| IMAGE_PAD, VISION_END = '<|image_pad|>', '<|vision_end|>' | |
| RESERVED = ['The document does not state this', 'The document gives conflicting answers'] | |
| NOT_STATED, CONFLICTING = 'not_stated', 'conflicting' | |
| SINGLE_R = ('Task: choose the single option that the document establishes as the answer. If the document does not establish any of ' | |
| 'the other listed answers, choose the option that says it does not state this. If the document establishes two different ' | |
| 'listed answers and does not say which prevails, choose the option that says it gives conflicting answers. A replaced or ' | |
| 'withdrawn statement establishes nothing. Respond with exactly one letter. Do not explain.') | |
| ORDERED_R = ('Task: the options form an ordered scale, lowest first, followed by two special options. Choose the single level that the ' | |
| 'document establishes. If the document does not establish any level, choose the option that says it does not state this. If ' | |
| 'it establishes two different levels and does not say which prevails, choose the option that says it gives conflicting ' | |
| 'answers. A replaced or withdrawn statement establishes nothing. Respond with exactly one letter. Do not explain.') | |
| SINGLE_S = CHOICE_TASK | |
| ORDERED_S = ('Task: the options form an ordered scale, lowest first. Choose the single level that the document best supports. ' | |
| 'Respond with exactly one letter. Do not explain.') | |
| SUFFICIENCY = ('Task: classify what the document establishes about the answer to the question, using only the document. ' | |
| 'A = it establishes exactly one of the listed answers. B = it does not establish any of the listed answers. ' | |
| 'C = it establishes two or more different listed answers and does not say which prevails. ' | |
| 'A missing fact is not a negative fact. A replaced or withdrawn statement establishes nothing. ' | |
| 'Respond with exactly one letter: A, B, or C. Do not explain.') | |
| LABEL = ('Task: decide whether the label applies, using only the document. A = the document establishes that it applies, and nothing in ' | |
| 'it establishes that it does not. B = the document establishes that it does not apply, and nothing in it establishes that it ' | |
| 'does. C = the document establishes neither. D = the document establishes both. A missing fact is not a negative fact. ' | |
| 'Evidence about another person or subject does not count. A replaced or withdrawn statement establishes nothing. ' | |
| 'Respond with exactly one letter: A, B, C, or D. Do not explain.') | |
| OPTION = ('Task: decide whether the proposed answer is correct, using only the document. A = the document establishes this answer, and ' | |
| 'nothing in it establishes a different one. B = the document establishes a different answer, or establishes that this one is ' | |
| 'wrong, and nothing in it establishes this one. C = the document establishes neither. D = the document establishes both this ' | |
| 'answer and a different one. A missing fact is not a negative fact. A replaced or withdrawn statement establishes nothing. ' | |
| 'Respond with exactly one letter: A, B, C, or D. Do not explain.') | |
| def _lettered(options): | |
| return '\n'.join(f'{LETTERS[i]}. {text}' for i, text in enumerate(options)) | |
| def listwise_block(question, options, ordered=False, reserved=True): | |
| """Caller options in caller order; with `reserved` the service appends the two reserved outcomes.""" | |
| if not 2 <= len(options) <= 8 or len(set(options)) != len(options): | |
| raise ValueError('needs 2 to 8 distinct options') | |
| shown = list(options) + (RESERVED if reserved else []) | |
| task = (ORDERED_R if ordered else SINGLE_R) if reserved else (ORDERED_S if ordered else SINGLE_S) | |
| return task + '\nQuestion: ' + question + '\nOptions:\n' + _lettered(shown) + '\nAnswer (one letter):', len(shown) | |
| def sufficiency_block(question, options): | |
| return SUFFICIENCY + '\nQuestion: ' + question + '\nListed answers:\n' + '\n'.join('- ' + o for o in options) + '\nAnswer (one letter):', 3 | |
| def label_block(question, label): | |
| return LABEL + '\nQuestion: ' + question + '\nLabel: ' + label + '\nAnswer (one letter):', 4 | |
| def option_block(question, option): | |
| return OPTION + '\nQuestion: ' + question + '\nProposed answer: ' + option + '\nAnswer (one letter):', 4 | |
| def four_state_block(question): | |
| return boolean_block(question), 4 | |
| CTL = {'start': None} # None: adapter off. int s: adapter applied to positions >= s of the current call. | |
| def _switch_lora_class(): | |
| import mlx.core as mx | |
| import mlx.nn as nn | |
| class SwitchLoRA(nn.Module): | |
| def __init__(self, linear, lora_a, lora_b, scale): | |
| super().__init__() | |
| self.linear = linear | |
| self.lora_a = lora_a | |
| self.lora_b = lora_b | |
| self.scale = scale | |
| def __call__(self, x): | |
| y = self.linear(x) | |
| s = CTL['start'] | |
| if s is None or s >= x.shape[1]: | |
| return y | |
| xs = x if s <= 0 else x[:, s:] | |
| z = (self.scale * ((xs @ self.lora_a) @ self.lora_b)).astype(y.dtype) | |
| if s <= 0: | |
| return y + z | |
| return mx.concatenate([y[:, :s], y[:, s:] + z], axis=1) | |
| return SwitchLoRA | |
| def attach_adapter(lm, path, scale=2.0): | |
| """Wrap the projections named in the adapter file. Returns (wrapped count, rank).""" | |
| import mlx.core as mx | |
| weights = mx.load(str(path)) | |
| cls = _switch_lora_class() | |
| names = sorted({k.rsplit('.', 1)[0] for k in weights}) | |
| rank = None | |
| for name in names: | |
| parts = name.split('.') # model.layers.N.owner.key | |
| owner = lm.model.layers[int(parts[2])] | |
| for p in parts[3:-1]: | |
| owner = getattr(owner, p) | |
| a, b = weights[name + '.lora_a'].astype(mx.float32), weights[name + '.lora_b'].astype(mx.float32) | |
| rank = a.shape[1] | |
| setattr(owner, parts[-1], cls(getattr(owner, parts[-1]), a, b, scale)) | |
| lm.freeze() | |
| mx.eval(lm.parameters()) | |
| return len(names), rank | |
| def as_parts(document): | |
| """A document is a string (text) or a list of parts {'text': str} / {'image': path}.""" | |
| if isinstance(document, str): | |
| return [{'text': document}] | |
| if not document or any(set(p) not in ({'text'}, {'image'}) for p in document): | |
| raise ValueError('parts must be {"text": ...} or {"image": ...}') | |
| return document | |
| def document_key(document): | |
| h = hashlib.sha256() | |
| for p in as_parts(document): | |
| h.update(b'T' + p['text'].encode() if 'text' in p else b'I' + bytes.fromhex(sha(p['image']))) | |
| return h.hexdigest() | |
| class ContractEngine(Engine): | |
| def __init__(self, arithmetic='float32', ledger=None, adapter=None, placement='question'): | |
| if placement not in ('question', 'full'): | |
| raise ValueError('placement must be question or full') | |
| super().__init__(arithmetic, ledger, adapter=None) | |
| self.placement = placement | |
| self.adapter = None | |
| self.rank = None | |
| if adapter is not None: | |
| self.wrapped, self.rank = attach_adapter(self.lm, adapter) | |
| self.adapter = str(adapter) | |
| self._pad = self.t.convert_tokens_to_ids(IMAGE_PAD) | |
| # ---- rendering ------------------------------------------------------------------- | |
| def _content(self, parts): | |
| out = '' | |
| for i, p in enumerate(parts): | |
| if 'text' in p: | |
| out += ('\n' if i and 'image' in parts[i - 1] else '') + p['text'] | |
| else: | |
| out += ('\n' if i and 'text' in parts[i - 1] else '') + PAGE | |
| return out | |
| def _text(self, parts, block): | |
| msgs = [{'role': 'system', 'content': SYSTEM}, {'role': 'user', 'content': 'Document:\n' + self._content(parts) + '\n\n' + block}] | |
| return self.t.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True, enable_thinking=False) | |
| def _expand(self, ids, counts): | |
| out, k = [], 0 | |
| for tok in ids: | |
| if tok == self._pad: | |
| out += [tok] * counts[k] | |
| k += 1 | |
| else: | |
| out.append(tok) | |
| if k != len(counts): | |
| raise ValueError('image placeholder count differs from the number of images') | |
| return out | |
| def _pixels(self, parts): | |
| from PIL import Image | |
| images = [Image.open(p['image']).convert('RGB') for p in parts if 'image' in p] | |
| if not images: | |
| return None, None, [] | |
| out = self.processor.image_processor(images=images) | |
| grid = np.array(out['image_grid_thw']) | |
| merge = self.model.config.vision_config.spatial_merge_size ** 2 | |
| counts = [int(g[0] * g[1] * g[2]) // merge for g in grid] | |
| return self.mx.array(np.array(out['pixel_values'])), self.mx.array(grid), counts | |
| def _vision(self, pixels, grid): | |
| """Vision tower one image at a time (memory stays flat for long scans).""" | |
| mx = self.mx | |
| dtype = self.model.vision_tower.patch_embed.proj.weight.dtype | |
| feats, offset = [], 0 | |
| for i in range(grid.shape[0]): | |
| n = int(np.prod(np.array(grid[i]))) | |
| h, _ = self.model.vision_tower(pixels[offset:offset + n].astype(dtype), grid[i:i + 1]) | |
| mx.eval(h) | |
| feats.append(h) | |
| offset += n | |
| return mx.concatenate(feats, axis=0) | |
| # ---- execution ------------------------------------------------------------------- | |
| def _forward(self, ids, pos, cache, embeds=None, chunk=CHUNK, adapter_from=None): | |
| """Run ids through the language model; logits for the final position. adapter_from indexes into ids.""" | |
| mx = self.mx | |
| n = len(ids) | |
| if self.ledger is not None: | |
| self.ledger.add(forwards=1, input_tokens=n) | |
| step = n if chunk is None else chunk | |
| if cache is None and step < n: | |
| raise ValueError('Chunked execution needs a cache') | |
| logits = None | |
| for a in range(0, n, step): | |
| b = min(a + step, n) | |
| CTL['start'] = None if adapter_from is None else max(0, adapter_from - a) | |
| kwargs = dict(cache=cache, position_ids=pos[:, :, a:b], skip_logits=True) | |
| if embeds is not None: | |
| kwargs['inputs_embeds'] = embeds[:, a:b] | |
| if b == n: | |
| out = self.lm(mx.array([ids[a:b]]), return_hidden=True, **kwargs) | |
| lg = self.lm.lm_head(out.hidden_states[-1][:, -1:, :])[0, -1].astype(mx.float32) | |
| mx.eval(lg) | |
| logits = np.array(lg) | |
| del out | |
| else: | |
| self.lm(mx.array([ids[a:b]]), **kwargs) | |
| mx.eval([c.state for c in cache]) | |
| CTL['start'] = None | |
| if not np.isfinite(logits).all(): | |
| raise ValueError('Non-finite model output') | |
| if mx.get_peak_memory() > 100 * 2**30: | |
| raise MemoryError('100 GiB MLX allocation cap exceeded') | |
| return logits | |
| def _positions(self, start, n): | |
| mx = self.mx | |
| return mx.broadcast_to(mx.arange(start, start + n)[None, None, :], (3, 1, n)) | |
| def prefill(self, document, chunk=CHUNK): | |
| """Encode the task-neutral document prefix once. Returns state format v3.""" | |
| mx = self.mx | |
| parts = as_parts(document) | |
| text = self._text(parts, 'X') | |
| end = text.rfind('\n\nX') | |
| if end < 0: | |
| raise ValueError('Document boundary missing') | |
| raw = self.t.encode(text[:end], add_special_tokens=False) | |
| if 'text' in parts[-1]: | |
| raw = raw[:-1] # leave the boundary token uncached: its tokenisation can depend on what follows | |
| pixels, grid, counts = self._pixels(parts) | |
| ids = self._expand(raw, counts) | |
| if not ids or len(ids) > TOKEN_CAP: | |
| raise ValueError(f'prefix of {len(ids)} tokens is empty or exceeds the {TOKEN_CAP}-token cap') | |
| embeds, delta, feats = None, 0, None | |
| if counts: | |
| feats = self._vision(pixels, grid) | |
| f = self.model.get_input_embeddings(mx.array([ids]), pixels, image_grid_thw=grid, cached_image_features=feats) | |
| embeds, pos = f.inputs_embeds, f.position_ids | |
| delta = int(np.array(f.rope_deltas).reshape(-1)[0]) | |
| if delta != int(np.array(pos).max()) + 1 - len(ids): | |
| raise ValueError('rope delta disagrees with the prefix positions') | |
| else: | |
| pos = self._positions(0, len(ids)) | |
| cache = self.lm.make_cache() | |
| self._forward(ids, pos, cache, embeds, chunk, adapter_from=0 if (self.adapter and self.placement == 'full') else None) | |
| mx.clear_cache() | |
| return {'format': STATE_FORMAT, 'prefix_ids': ids, 'cache': cache, 'rope_delta': delta, 'parts': parts, | |
| 'image_tokens': int(sum(counts)), 'image_sha256': [sha(p['image']) for p in parts if 'image' in p], | |
| 'layout': ''.join('T' if 'text' in p else 'I' for p in parts), | |
| 'prefix_sha256': hashlib.sha256(json.dumps(ids).encode()).hexdigest(), | |
| '_pixels': pixels, '_grid': grid, '_counts': counts, '_feats': feats} | |
| def letters_for(self, text, n): | |
| ids = self.t.encode(text, add_special_tokens=False) | |
| out = [] | |
| for letter in LETTERS[:n]: | |
| if letter not in self._letters: | |
| ext = self.t.encode(text + letter, add_special_tokens=False) | |
| if len(ext) != len(ids) + 1 or ext[:-1] != ids: | |
| raise ValueError('Unstable answer-letter continuation') | |
| self._letters[letter] = ext[-1] | |
| out.append(self._letters[letter]) | |
| return out, ids | |
| def ask(self, state, block, n_letters, execution='cached', full_chunk=None): | |
| """Answer-letter distribution for one task block against a prefilled state.""" | |
| mx = self.mx | |
| text = self._text(state['parts'], block) | |
| letters, raw = self.letters_for(text, n_letters) | |
| ids = self._expand(raw, state['_counts']) | |
| prefix = state['prefix_ids'] | |
| P = len(prefix) | |
| fallback = None | |
| if ids[:P] != prefix: | |
| execution, fallback = 'full', 'prefix token mismatch' | |
| on = self.adapter is not None | |
| if execution == 'cached': | |
| branch = copy.deepcopy(state['cache']) | |
| logits = self._forward(ids[P:], self._positions(P + state['rope_delta'], len(ids) - P), branch, chunk=None, | |
| adapter_from=0 if on else None) | |
| del branch | |
| reused = P | |
| else: | |
| embeds = None | |
| if state['_counts']: | |
| f = self.model.get_input_embeddings(mx.array([ids]), state['_pixels'], image_grid_thw=state['_grid'], | |
| cached_image_features=None if len(state['_counts']) <= 6 else state['_feats']) | |
| embeds, pos = f.inputs_embeds, f.position_ids | |
| else: | |
| pos = self._positions(0, len(ids)) | |
| start = None if not on else (0 if self.placement == 'full' else P) | |
| logits = self._forward(ids, pos, self.lm.make_cache() if full_chunk else None, embeds, full_chunk, adapter_from=start) | |
| reused = 0 | |
| full = softmax(logits) | |
| mx.clear_cache() | |
| return {'letter_logits': logits[letters].astype(np.float64), 'probabilities': softmax(logits[letters]), | |
| 'mass': float(full[letters].sum()), 'top_is_letter': int(logits.argmax()) in letters, 'execution': execution, | |
| 'fallback': fallback, 'prompt_tokens': len(ids), 'reused_prefix_tokens': reused, 'branch_tokens': len(ids) - reused} | |
| def identity(engine): | |
| """Runtime fingerprint for state format v3. With question placement the adapter is not part of the state.""" | |
| from importlib.metadata import version | |
| ip = engine.processor.image_processor | |
| payload = {'contract': CONTRACT, 'state_format': STATE_FORMAT, 'system_sha256': hashlib.sha256(SYSTEM.encode()).hexdigest(), | |
| 'arithmetic': engine.arithmetic, 'engine_sha256': sha(Path(__file__)), 'base_engine_sha256': sha(ROOT / 'solomon/engine_reference.py'), | |
| 'model_config_sha256': sha(MODEL / 'config.json'), | |
| 'image_processor': {'class': type(ip).__name__, 'config_sha256': sha(MODEL / 'preprocessor_config.json') | |
| if (MODEL / 'preprocessor_config.json').exists() else None}, | |
| 'state_adapter': (engine.adapter and sha(engine.adapter)) if engine.placement == 'full' else None, | |
| 'versions': {p: version(p) for p in ('mlx', 'mlx-vlm', 'transformers')}} | |
| payload['fingerprint'] = hashlib.sha256(json.dumps(payload, sort_keys=True).encode()).hexdigest() | |
| payload['branch_adapter'] = engine.adapter and {'sha256': sha(engine.adapter), 'placement': engine.placement, 'rank': engine.rank} | |
| return payload | |