mlx-decision 0.3.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (44) hide show
  1. mlx_decision/__init__.py +32 -0
  2. mlx_decision/answers.py +50 -0
  3. mlx_decision/backbones/__init__.py +0 -0
  4. mlx_decision/backbones/qwen3_5/__init__.py +3 -0
  5. mlx_decision/backbones/qwen3_5/activations.py +74 -0
  6. mlx_decision/backbones/qwen3_5/base.py +143 -0
  7. mlx_decision/backbones/qwen3_5/gated_delta.py +651 -0
  8. mlx_decision/backbones/qwen3_5/load.py +210 -0
  9. mlx_decision/backbones/qwen3_5/mrope.py +40 -0
  10. mlx_decision/backbones/qwen3_5/positions.py +54 -0
  11. mlx_decision/backbones/qwen3_5/preprocess.py +119 -0
  12. mlx_decision/backbones/qwen3_5/qwen3_5.py +336 -0
  13. mlx_decision/backbones/qwen3_5/qwen3_next.py +165 -0
  14. mlx_decision/backbones/qwen3_5/rope_utils.py +377 -0
  15. mlx_decision/backbones/qwen3_5/vision.py +216 -0
  16. mlx_decision/backend.py +36 -0
  17. mlx_decision/benchmark.py +152 -0
  18. mlx_decision/calibration.py +194 -0
  19. mlx_decision/cli.py +681 -0
  20. mlx_decision/convert.py +54 -0
  21. mlx_decision/download.py +43 -0
  22. mlx_decision/errors.py +31 -0
  23. mlx_decision/hub.py +51 -0
  24. mlx_decision/images.py +122 -0
  25. mlx_decision/interactive.py +478 -0
  26. mlx_decision/mixed.py +164 -0
  27. mlx_decision/model.py +126 -0
  28. mlx_decision/models/__init__.py +0 -0
  29. mlx_decision/models/clef/__init__.py +0 -0
  30. mlx_decision/models/clef/convert.py +154 -0
  31. mlx_decision/models/clef/encode.py +171 -0
  32. mlx_decision/models/clef/head.py +205 -0
  33. mlx_decision/models/clef/model.py +204 -0
  34. mlx_decision/registry.py +85 -0
  35. mlx_decision/server.py +148 -0
  36. mlx_decision/types.py +204 -0
  37. mlx_decision-0.3.0.dist-info/METADATA +529 -0
  38. mlx_decision-0.3.0.dist-info/RECORD +44 -0
  39. mlx_decision-0.3.0.dist-info/WHEEL +4 -0
  40. mlx_decision-0.3.0.dist-info/entry_points.txt +2 -0
  41. mlx_decision-0.3.0.dist-info/licenses/LICENSE +21 -0
  42. mlx_decision-0.3.0.dist-info/licenses/LICENSES/clef-Apache-2.0.txt +202 -0
  43. mlx_decision-0.3.0.dist-info/licenses/LICENSES/mlx-lm-MIT.txt +21 -0
  44. mlx_decision-0.3.0.dist-info/licenses/LICENSES/transformers-Apache-2.0.txt +202 -0
@@ -0,0 +1,32 @@
1
+ from .errors import DecisionError
2
+ from .model import DecisionModel, load
3
+ from .types import (
4
+ Choice,
5
+ ChoiceAnswer,
6
+ Noul,
7
+ NoulAnswer,
8
+ Request,
9
+ Response,
10
+ Result,
11
+ Score,
12
+ ScoreAnswer,
13
+ Usage,
14
+ )
15
+
16
+ __version__ = "0.3.0"
17
+
18
+ __all__ = [
19
+ "Choice",
20
+ "ChoiceAnswer",
21
+ "DecisionError",
22
+ "DecisionModel",
23
+ "Noul",
24
+ "NoulAnswer",
25
+ "Request",
26
+ "Response",
27
+ "Result",
28
+ "Score",
29
+ "ScoreAnswer",
30
+ "Usage",
31
+ "load",
32
+ ]
@@ -0,0 +1,50 @@
1
+ """Turn per-option probabilities into answers. Shared by every backend."""
2
+
3
+ from collections.abc import Mapping, Sequence
4
+
5
+ from .types import Answer, Choice, ChoiceAnswer, Noul, NoulAnswer, Score, ScoreAnswer
6
+
7
+
8
+ def choice_confidence(probabilities: Sequence[float]) -> float:
9
+ """How far the top probability sits above an even split, from 0 to 1."""
10
+ n = len(probabilities)
11
+ if n < 2:
12
+ return 1.0
13
+ return max(0.0, (max(probabilities) - 1 / n) / (1 - 1 / n))
14
+
15
+
16
+ def score_confidence(probabilities: Sequence[float]) -> float:
17
+ """One minus the spread around the most likely level, relative to an even spread."""
18
+ n = len(probabilities)
19
+ if n < 2:
20
+ return 1.0
21
+ peak = max(range(n), key=probabilities.__getitem__)
22
+ spread = sum(p * abs(i - peak) for i, p in enumerate(probabilities))
23
+ even_spread = sum(abs(i - (n - 1) / 2) for i in range(n)) / n
24
+ return max(0.0, 1 - spread / even_spread)
25
+
26
+
27
+ def build_answer(question: Noul | Choice | Score, probabilities: Mapping[str, float]) -> Answer:
28
+ """Build the answer for one question from its option probabilities.
29
+
30
+ ``probabilities`` is keyed by option id: ``"true"``/``"false"`` for a noul,
31
+ the criteria keys for a choice, ``"0"``, ``"1"``, ... for a score.
32
+ """
33
+ if isinstance(question, Noul):
34
+ return NoulAnswer(noul=probabilities["true"])
35
+ if isinstance(question, Choice):
36
+ options = list(question.criteria)
37
+ ordered = {option: probabilities[option] for option in options}
38
+ return ChoiceAnswer(
39
+ choice=max(options, key=ordered.__getitem__),
40
+ confidence=choice_confidence(list(ordered.values())),
41
+ probabilities=ordered,
42
+ )
43
+ levels = [str(index) for index in range(len(question.criteria))]
44
+ values = [probabilities[level] for level in levels]
45
+ return ScoreAnswer(
46
+ score=sum(index * p for index, p in enumerate(values)),
47
+ confidence=score_confidence(values),
48
+ legend=dict(zip(levels, question.criteria, strict=True)),
49
+ probabilities=dict(zip(levels, values, strict=True)),
50
+ )
File without changes
@@ -0,0 +1,3 @@
1
+ from .qwen3_5 import Model, ModelArgs
2
+
3
+ __all__ = ["Model", "ModelArgs"]
@@ -0,0 +1,74 @@
1
+ # Copyright © 2023 Apple Inc.
2
+ #
3
+ # Copied from mlx-lm (https://github.com/ml-explore/mlx-lm), file
4
+ # mlx_lm/models/activations.py at commit 5cfec4cb39deba54210b3ff4d86f2337c7bc10b5,
5
+ # under the MIT License. See LICENSES/mlx-lm-MIT.txt.
6
+
7
+ from functools import partial
8
+
9
+ import mlx.core as mx
10
+ import mlx.nn as nn
11
+
12
+
13
+ @partial(mx.compile, shapeless=True)
14
+ def swiglu(gate, x):
15
+ return nn.silu(gate) * x
16
+
17
+
18
+ @partial(mx.compile, shapeless=True)
19
+ def precise_swiglu(h, gate, x):
20
+ gate = nn.silu(gate.astype(mx.float32))
21
+ x = x.astype(mx.float32)
22
+ return (gate * x).astype(h.dtype)
23
+
24
+
25
+ @partial(mx.compile, shapeless=True)
26
+ def swiglu_oai(
27
+ gate: mx.array, x: mx.array, alpha: float = 1.702, limit: float = 7.0
28
+ ) -> mx.array:
29
+ # GPT-OSS variant: clip the gate from above only, and add 1 to the linear part.
30
+ gate = mx.minimum(gate, limit)
31
+ x = mx.clip(x, -limit, limit)
32
+ return (x + 1.0) * (gate * mx.sigmoid(alpha * gate))
33
+
34
+
35
+ class SwigluOAI(nn.Module):
36
+ def __init__(self, alpha: float = 1.702, limit: float = 7.0):
37
+ super().__init__()
38
+ self._alpha = alpha
39
+ self._limit = limit
40
+
41
+ def __call__(self, x: mx.array, gate: mx.array) -> mx.array:
42
+ return swiglu_oai(gate, x, alpha=self._alpha, limit=self._limit)
43
+
44
+
45
+ @partial(mx.compile, shapeless=True)
46
+ def xielu(x, alpha_p, alpha_n, beta, eps):
47
+ alpha_p = nn.softplus(alpha_p)
48
+ alpha_n = beta + nn.softplus(alpha_n)
49
+ return mx.where(
50
+ x > 0,
51
+ alpha_p * mx.square(x) + beta * x,
52
+ (mx.expm1(mx.minimum(x, eps)) - x) * alpha_n + beta * x,
53
+ )
54
+
55
+
56
+ class XieLU(nn.Module):
57
+ def __init__(
58
+ self,
59
+ alpha_p_init=0.8,
60
+ alpha_n_init=0.8,
61
+ beta=0.5,
62
+ eps=-1e-6,
63
+ ):
64
+ super().__init__()
65
+ alpha_p_tensor = mx.array(alpha_p_init)
66
+ alpha_n_tensor = mx.array(alpha_n_init - beta)
67
+ self.alpha_p = mx.log(mx.exp(alpha_p_tensor) - 1)
68
+ self.alpha_n = mx.log(mx.exp(alpha_n_tensor) - 1)
69
+
70
+ self.beta = mx.array(beta)
71
+ self.eps = mx.array(eps)
72
+
73
+ def __call__(self, x: mx.array) -> mx.array:
74
+ return xielu(x, self.alpha_p, self.alpha_n, self.beta, self.eps)
@@ -0,0 +1,143 @@
1
+ # Copyright © 2023 Apple Inc.
2
+ #
3
+ # Copied from mlx-lm (https://github.com/ml-explore/mlx-lm), file
4
+ # mlx_lm/models/base.py at commit 5cfec4cb39deba54210b3ff4d86f2337c7bc10b5,
5
+ # under the MIT License. See LICENSES/mlx-lm-MIT.txt.
6
+
7
+ import inspect
8
+ from dataclasses import dataclass
9
+ from typing import Optional
10
+
11
+ import mlx.core as mx
12
+ from mlx.utils import tree_map
13
+
14
+
15
+ @dataclass
16
+ class BaseModelArgs:
17
+ @classmethod
18
+ def from_dict(cls, params):
19
+ return cls(
20
+ **{
21
+ k: v
22
+ for k, v in params.items()
23
+ if k in inspect.signature(cls).parameters
24
+ }
25
+ )
26
+
27
+
28
+ def create_causal_mask(
29
+ N: int,
30
+ offset: int = 0,
31
+ window_size: Optional[int] = None,
32
+ right_padding: Optional[mx.array] = None,
33
+ left_padding: Optional[mx.array] = None,
34
+ ):
35
+ rinds = mx.arange(offset + N)
36
+ linds = mx.arange(offset, offset + N) if offset else rinds
37
+ linds = linds[:, None]
38
+ rinds = rinds[None]
39
+ mask = linds >= rinds
40
+ if window_size is not None:
41
+ mask = mask & (linds < rinds + window_size)
42
+ if right_padding is not None:
43
+ mask = mask & (rinds < mx.expand_dims((offset + N) - right_padding, (1, 2, 3)))
44
+ if left_padding is not None:
45
+ mask = mask & (mx.expand_dims(left_padding, (1, 2, 3)) <= rinds)
46
+ return mask
47
+
48
+
49
+ def create_attention_mask(
50
+ h, cache=None, window_size: Optional[int] = None, return_array: bool = False
51
+ ):
52
+ N = h.shape[1]
53
+ if cache and hasattr(cache, "make_mask"):
54
+ return cache.make_mask(N, return_array=return_array, window_size=window_size)
55
+ if N == 1:
56
+ return None
57
+ if return_array or (window_size and N > window_size):
58
+ return create_causal_mask(N, window_size=window_size)
59
+ return "causal"
60
+
61
+
62
+ def create_ssm_mask(h, cache=None):
63
+ if cache and hasattr(cache, "make_mask"):
64
+ return cache.make_mask(h.shape[1])
65
+ return None
66
+
67
+
68
+ def quantized_scaled_dot_product_attention(
69
+ queries: mx.array,
70
+ q_keys: tuple[mx.array, mx.array, mx.array],
71
+ q_values: tuple[mx.array, mx.array, mx.array],
72
+ scale: float,
73
+ mask: Optional[mx.array],
74
+ group_size: int = 64,
75
+ bits: int = 8,
76
+ ) -> mx.array:
77
+ B, n_q_heads, L, D = queries.shape
78
+ n_kv_heads = q_keys[0].shape[-3]
79
+ n_repeats = n_q_heads // n_kv_heads
80
+
81
+ queries *= scale
82
+
83
+ if n_repeats > 1:
84
+ queries = mx.reshape(queries, (B, n_kv_heads, n_repeats, L, D))
85
+ q_keys = tree_map(lambda x: mx.expand_dims(x, axis=-3), q_keys)
86
+ q_values = tree_map(lambda x: mx.expand_dims(x, axis=-3), q_values)
87
+
88
+ scores = mx.quantized_matmul(
89
+ queries, *q_keys, transpose=True, group_size=group_size, bits=bits
90
+ )
91
+ if mask is not None:
92
+ if isinstance(mask, str):
93
+ qL, kL = scores.shape[-2:]
94
+ q_indices = mx.arange(kL - qL, kL)
95
+ k_indices = mx.arange(kL)
96
+ mask = q_indices[:, None] >= k_indices[None]
97
+ if n_repeats > 1 and mask.ndim > 3:
98
+ mask = mx.expand_dims(mask, -3)
99
+ if mask.dtype == mx.bool_:
100
+ scores = mx.where(mask, scores, mx.finfo(scores.dtype).min)
101
+ else:
102
+ scores += mask
103
+ scores = mx.softmax(scores, axis=-1, precise=True)
104
+ out = mx.quantized_matmul(
105
+ scores, *q_values, transpose=False, group_size=group_size, bits=bits
106
+ )
107
+
108
+ if n_repeats > 1:
109
+ out = mx.reshape(out, (B, n_q_heads, L, D))
110
+
111
+ return out
112
+
113
+
114
+ def scaled_dot_product_attention(
115
+ queries,
116
+ keys,
117
+ values,
118
+ cache,
119
+ scale: float,
120
+ mask: Optional[mx.array],
121
+ sinks: Optional[mx.array] = None,
122
+ ) -> mx.array:
123
+ if hasattr(cache, "bits"):
124
+ if sinks is not None:
125
+ raise ValueError("Quantized SDPA does not support attention sinks.")
126
+ return quantized_scaled_dot_product_attention(
127
+ queries,
128
+ keys,
129
+ values,
130
+ scale=scale,
131
+ mask=mask,
132
+ group_size=cache.group_size,
133
+ bits=cache.bits,
134
+ )
135
+ else:
136
+ return mx.fast.scaled_dot_product_attention(
137
+ queries,
138
+ keys,
139
+ values,
140
+ scale=scale,
141
+ mask=mask,
142
+ sinks=sinks,
143
+ )