adaptible 1.0.0a1__tar.gz

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.
@@ -0,0 +1,4 @@
1
+ __pycache__
2
+ uv.lock
3
+ .venv
4
+ .gitignore
@@ -0,0 +1,17 @@
1
+ Metadata-Version: 2.4
2
+ Name: adaptible
3
+ Version: 1.0.0a1
4
+ Summary: Stateful LLM serving instances that self-reflect and learn from their mistakes during periods of low server utilization.
5
+ Requires-Python: >=3.13
6
+ Requires-Dist: absl-py
7
+ Requires-Dist: fastapi[all]
8
+ Requires-Dist: immutabledict
9
+ Requires-Dist: lm-eval
10
+ Requires-Dist: mlx
11
+ Requires-Dist: mlx-lm
12
+ Requires-Dist: optax
13
+ Requires-Dist: torch
14
+ Requires-Dist: tqdm
15
+ Requires-Dist: transformers
16
+ Requires-Dist: uvicorn[standard]
17
+ Requires-Dist: vizible
File without changes
@@ -0,0 +1,3 @@
1
+ """Adaptible - LLMs that can wander."""
2
+
3
+ from ._src import app, InteractionHistory, StatefulLLM, MutableHostedLLM
@@ -0,0 +1,6 @@
1
+ """Adaptible - LLMs that can wander."""
2
+
3
+ from ._api import app
4
+ from ._classes import InteractionHistory
5
+ from ._llm import StatefulLLM
6
+ from ._server import MutableHostedLLM
@@ -0,0 +1,142 @@
1
+ """Standard Interaction Logic for a Stateful LLM"""
2
+
3
+ import asyncio
4
+ from asyncio import log
5
+ import collections
6
+ import os
7
+ import time
8
+ import threading
9
+ from typing import Any, List
10
+
11
+
12
+ from fastapi import FastAPI, HTTPException
13
+ from fastapi.staticfiles import StaticFiles
14
+ from fastapi.responses import StreamingResponse
15
+ import tqdm
16
+ import vizible
17
+
18
+ from ._classes import (
19
+ InteractionHistory,
20
+ InteractionRequest,
21
+ InteractionResponse,
22
+ ReviewResponse,
23
+ SyncResponse,
24
+ )
25
+ from ._llm import StatefulLLM
26
+
27
+ app = FastAPI(
28
+ title="Stateful Self-Improving LLM Server",
29
+ description="An API for a stateful LLM that tries to improve itself over time.",
30
+ )
31
+ app.mount(
32
+ "/static",
33
+ StaticFiles(
34
+ directory=os.path.join(os.path.dirname(__file__), "static"),
35
+ html=True,
36
+ ),
37
+ name="static",
38
+ )
39
+
40
+ # In-memory store or interaction history from the current session.
41
+ interaction_history: List[InteractionHistory] = []
42
+ unreviewed_interaction_history_indices: List[int] = []
43
+
44
+ outstanding_tasks: collections.deque[Any] = collections.deque([])
45
+
46
+ # Instantiate the stateful model. This is a singleton for the lifecycle of the app.
47
+ model = StatefulLLM()
48
+
49
+
50
+ @app.post("/interact", response_model=InteractionResponse)
51
+ async def interact_with_model(request: InteractionRequest):
52
+ """
53
+ Main endpoint for interacting with the LLM (Forward-pass).
54
+ """
55
+ if not request.prompt:
56
+ raise HTTPException(status_code=400, detail="Prompt cannot be empty.")
57
+ print(f"Received request: {request}")
58
+ # Generate the response using the current state of the model.
59
+ response_text = model.generate_response(request.prompt)
60
+
61
+ # Store the interaction for later review
62
+ interaction_idx = len(interaction_history)
63
+ unreviewed_interaction_history_indices.append(len(interaction_history))
64
+ interaction_history.append(
65
+ InteractionHistory(
66
+ idx=interaction_idx,
67
+ user_input=request.prompt,
68
+ llm_response=response_text,
69
+ reviewed=False,
70
+ timestamp=time.time(),
71
+ )
72
+ )
73
+ return {"response": response_text, "interaction_idx": interaction_idx}
74
+
75
+
76
+ @app.post("/stream_interact")
77
+ async def stream_interact_with_model(request: InteractionRequest):
78
+ """Stream interaction with model."""
79
+ # TODO - enable logging.
80
+ return StreamingResponse(
81
+ model.stream_response(request.prompt), media_type="text/plain"
82
+ )
83
+
84
+ @app.post("/trigger_review", response_model=ReviewResponse)
85
+ async def trigger_review_cycle():
86
+ """Manually triggers the self-correction and training cycle."""
87
+ unreviewed_count = len(unreviewed_interaction_history_indices)
88
+ if unreviewed_count == 0:
89
+ return {
90
+ "message": "No unreviewed interactions to process.",
91
+ "unreviewed_count": unreviewed_count,
92
+ }
93
+ unreviewed_interaction_history = [
94
+ interaction_history[idx] for idx in unreviewed_interaction_history_indices
95
+ ]
96
+ outstanding_tasks.append(
97
+ asyncio.to_thread(model.self_correct_and_train, unreviewed_interaction_history)
98
+ )
99
+ return {
100
+ "message": "Self-correction and training cycle has been initiated in the background.",
101
+ "unreviewed_count": unreviewed_count,
102
+ }
103
+
104
+
105
+ @app.get("/sync", response_model=SyncResponse)
106
+ async def sync_server_for_background_tasks():
107
+ """Waits for any tasks to background complete."""
108
+ start_time = time.time()
109
+ num_tasks = len(outstanding_tasks)
110
+ print("Waiting for model state to stabilize...")
111
+ lock = threading.Lock()
112
+ with (
113
+ lock,
114
+ tqdm.tqdm(desc="Waiting for server to sync.", unit=" Seconds") as server_pbar,
115
+ ):
116
+ vizible.green("Finished background tasks")
117
+ while not model.ok:
118
+ log.logger.info(
119
+ "Waiting for model server to sync. Is model is %s ok.",
120
+ "" if model.ok else "not ",
121
+ )
122
+ server_pbar.update(1)
123
+ await asyncio.sleep(1)
124
+ vizible.green("Model has reached a stable state!!!")
125
+ elapsed_time = time.time() - start_time
126
+ return {
127
+ "message": "Sync'd all background tasks.",
128
+ "tasks_count": num_tasks,
129
+ "elapsed_time": elapsed_time,
130
+ }
131
+
132
+
133
+ @app.get("/history")
134
+ async def get_history():
135
+ """Returns the full interaction history."""
136
+ return {"history": interaction_history}
137
+
138
+
139
+ @app.get("/status")
140
+ async def check_is_running():
141
+ """Basic response to signal that the server is operational."""
142
+ return {"status": "up"}
@@ -0,0 +1,80 @@
1
+ """Common class definitions."""
2
+
3
+ import dataclasses
4
+ import mlx.core as mx
5
+ from pydantic import BaseModel
6
+
7
+
8
+ @dataclasses.dataclass
9
+ class InteractionHistory:
10
+ """Event turn during user-LLM dialog
11
+
12
+ Attributes:
13
+ idx: Index of current turn amongst all global turns.
14
+ user_input: User-provided prompt.
15
+ llm_response: LLM response.
16
+ reviewed: Whether this interaction has been reviewed already.
17
+ timestamp: When the interaction took place, measured in seconds.
18
+ """
19
+ idx: int
20
+ user_input: str
21
+ llm_response: str = ""
22
+ reviewed: bool = False
23
+ timestamp: float = 0.0
24
+
25
+ @dataclasses.dataclass
26
+ class TrainingExample:
27
+ """Pre-tokenized training data
28
+
29
+ Attributes:
30
+ user_input: Actual model response.
31
+ label: Target model response.
32
+ mask: Mask of model response that dictates which parts are used in training.
33
+ """
34
+ input: mx.array
35
+ label: mx.array
36
+ mask: mx.array
37
+
38
+
39
+ class InteractionRequest(BaseModel):
40
+ """User prompt to be sent to the LLM.
41
+
42
+ Attributes:
43
+ prompt: User-provided input.
44
+ """
45
+ prompt: str
46
+
47
+
48
+ class InteractionResponse(BaseModel):
49
+ """LLM response to user-provided prompt
50
+
51
+ Attributes:
52
+ response: LLM-generated text response.
53
+ interaction_id: Index of the current response within the context of the current session.
54
+ """
55
+ response: str
56
+ interaction_idx: int
57
+
58
+
59
+ class ReviewResponse(BaseModel):
60
+ """Response to initiating asynchronous review of entire unreviewed interaction history.
61
+
62
+ Attributes:
63
+ message: Human-readable output after completion of review.
64
+ unreviewed_count: the number of unreviewed interactions handled by this operation.
65
+ """
66
+ message: str
67
+ unreviewed_count: int
68
+
69
+
70
+ class SyncResponse(BaseModel):
71
+ """Response after server completes all unfinished background tasks.
72
+
73
+ Attributes:
74
+ message: Human-readable message.
75
+ tasks_count: Number of tasks waited on and successfully finished.
76
+ elapsed_time: Amount of time background tasks took to finish.
77
+ """
78
+ message: str
79
+ tasks_count: int
80
+ elapsed_time: float
@@ -0,0 +1,302 @@
1
+ """Stateful LLM."""
2
+
3
+ from typing import AsyncIterable, List, Tuple
4
+
5
+ import functools
6
+ import threading
7
+ import tqdm
8
+
9
+ import immutabledict
10
+ from mlx import optimizers
11
+ from mlx import nn
12
+ import mlx.core as mx
13
+ from mlx_lm.generate import generate, stream_generate
14
+ from mlx_lm.utils import load
15
+ from mlx_lm.tuner.utils import linear_to_lora_layers
16
+ from transformers.tokenization_utils import PreTrainedTokenizer
17
+ import vizible
18
+
19
+ from .libs import revise
20
+ from ._classes import InteractionHistory, TrainingExample
21
+
22
+
23
+ # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # #
24
+ # Default constants. #
25
+ # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # #
26
+ _MODEL_NAME = "lmstudio-community/DeepSeek-R1-0528-Qwen3-8B-MLX-4bit"
27
+ # _MODEL_NAME = "mlx-community/DeepSeek-R1-Distill-Qwen-1.5B-4bit"
28
+ _MAX_TOKENS = 32768
29
+ _LEARNING_RATE = 5e-5
30
+ _EPOCHS = 5
31
+ _NUM_LORA_LAYERS = 16
32
+ _LORA_PARAMETERS = immutabledict.immutabledict(
33
+ {"rank": 8, "dropout": 0.0, "scale": 20.0}
34
+ )
35
+ _USE_DORA = False
36
+
37
+
38
+ def _load(
39
+ model_name: str,
40
+ num_lora_layers: int,
41
+ use_dora: bool,
42
+ lora_parameters: dict | None = None,
43
+ ) -> Tuple[nn.Module, PreTrainedTokenizer]:
44
+ """Load model parameters and tokenizer.
45
+
46
+ Args:
47
+ model_name: Path or Huggingface name.
48
+ num_lora_layers: Number of LORA layers, if LORA is enabled.
49
+ use_dora: Whether to use DORA, if LORA is enabled.
50
+ lora_parameters: LORA hyperparameters. If not None, LORA will be enabled.
51
+
52
+ Returns:
53
+ Model and tokenizer.
54
+ """
55
+ model, wrapped_tokenizer = load(model_name)
56
+ print("Freezing all non-Lora model parameters.")
57
+ model.freeze()
58
+ if lora_parameters is not None:
59
+ linear_to_lora_layers(
60
+ model=model,
61
+ num_layers=num_lora_layers,
62
+ config=lora_parameters,
63
+ use_dora=use_dora,
64
+ )
65
+ return model, wrapped_tokenizer._tokenizer # pylint: disable=protected-access
66
+
67
+
68
+ def _loss_fn(
69
+ model: nn.Module,
70
+ inputs: mx.array,
71
+ targets: mx.array,
72
+ mask: mx.array,
73
+ ) -> mx.array:
74
+ logits: mx.array = model(inputs, mx.ones_like(inputs))
75
+ loss = nn.losses.cross_entropy(logits, targets, reduction="mean") * mask
76
+ normalized_loss = loss.sum() / mask.sum()
77
+ return normalized_loss
78
+
79
+
80
+ class StatefulLLM:
81
+ """Model container that bundles revision, learning, and serving logic."""
82
+
83
+ def __init__(
84
+ self,
85
+ model_name: str = _MODEL_NAME,
86
+ learning_rate: float = _LEARNING_RATE,
87
+ max_tokens: int = _MAX_TOKENS,
88
+ epochs: int = _EPOCHS,
89
+ num_lora_layers: int = _NUM_LORA_LAYERS,
90
+ lora_parameters: dict | None = _LORA_PARAMETERS,
91
+ use_dora: bool = _USE_DORA,
92
+ ) -> None:
93
+ """Initializes the StatefulLLM
94
+
95
+ Args:
96
+ model_name: Path or Huggingface name.
97
+ learning_rate: Backpropagation hyperparameter.
98
+ max_tokens: Maximum number of tokens to decode in a single turn.
99
+ epochs: Number of training epochs to perform on self-reflective model revisions.
100
+ num_lora_layers: Number of LORA layers, if LORA is enabled.
101
+ lora_parameters: LORA hyperparameters. If not None, LORA will be enabled.
102
+ use_dora: Whether to use DORA, if LORA is enabled.
103
+
104
+ Returns: None
105
+ """
106
+ self._lock = threading.Lock()
107
+ with self._lock:
108
+ self._messages: list[dict[str, str]] = []
109
+ self._model, self._tokenizer = _load(
110
+ model_name,
111
+ num_lora_layers=num_lora_layers,
112
+ lora_parameters=lora_parameters,
113
+ use_dora=use_dora,
114
+ )
115
+ self._model_name = model_name
116
+ self._optimizer = optimizers.AdamW(learning_rate=learning_rate)
117
+ self._epochs = epochs
118
+ self._max_tokens = max_tokens
119
+ self._bos_token = self._tokenizer.special_tokens_map.get("bos_token", "")
120
+ self._eos_token = self._tokenizer.special_tokens_map.get("eos_token", "")
121
+
122
+ self._model_is_stable = True
123
+ self._response_stream = None
124
+
125
+ @property
126
+ def ok(self) -> bool:
127
+ """Checks if the model is in a stable state (i.e. not in the middle of backprop)."""
128
+ return self._model_is_stable
129
+
130
+ def generate_response(
131
+ self,
132
+ prompt: str,
133
+ use_history: bool = True,
134
+ max_tokens: int | None = None,
135
+ ) -> str:
136
+ """Generates model response, including handling tokenization and de-tokenization.
137
+
138
+ Args:
139
+ prompt: User input.
140
+ use_history: Whether to include previous interaction history within the prompt.
141
+ max_tokens: Maximum number of tokens to decode in a single turn.
142
+
143
+ Returns:
144
+ Model-generated response.
145
+ """
146
+ vizible.blue(f"Generating response for {prompt = }")
147
+ message = {"role": "user", "content": prompt}
148
+ if max_tokens is None:
149
+ max_tokens = self._max_tokens
150
+ if use_history:
151
+ self._messages.append(message)
152
+ messages = self._messages
153
+ else:
154
+ messages = [message]
155
+ tokenized_prompt = self._tokenizer.apply_chat_template(messages)
156
+ model_response = generate(
157
+ self._model,
158
+ self._tokenizer,
159
+ prompt=tokenized_prompt,
160
+ verbose=True,
161
+ max_tokens=max_tokens,
162
+ )
163
+ return model_response.strip()
164
+
165
+ async def stream_response(
166
+ self,
167
+ prompt: str,
168
+ use_history: bool = True,
169
+ max_tokens: int | None = None,
170
+ ) -> AsyncIterable[str]:
171
+ """Stream generated response"""
172
+ message = {"role": "user", "content": prompt}
173
+ if max_tokens is None:
174
+ max_tokens = self._max_tokens
175
+ print(f'{self._messages = }')
176
+ if use_history:
177
+ self._messages.append(message)
178
+ messages = self._messages
179
+ else:
180
+ messages = [message]
181
+ tokenized_prompt = self._tokenizer.apply_chat_template(messages)
182
+ responses = []
183
+ try:
184
+ for response in stream_generate(
185
+ self._model,
186
+ self._tokenizer,
187
+ prompt=tokenized_prompt,
188
+ max_tokens=max_tokens,
189
+ ):
190
+ text = response.text
191
+ responses.append(text)
192
+ yield text
193
+ finally:
194
+ self._response_stream = "".join(responses)
195
+
196
+ def _tokenize(self, inp: str, dtype: mx.Dtype = mx.int32) -> mx.array:
197
+ return mx.array(
198
+ self._tokenizer.encode(inp, add_special_tokens=False), dtype=dtype
199
+ )
200
+
201
+ def _self_correct(
202
+ self,
203
+ interaction_history: List[InteractionHistory],
204
+ indices_to_review: List[int] | None,
205
+ verbose: bool,
206
+ ) -> TrainingExample:
207
+ self._model_is_stable = False
208
+ vizible.green("\n--- Starting Self-Correction and Training Cycle ---")
209
+ if indices_to_review is None:
210
+ indices_to_review = list(range(len(interaction_history)))
211
+ if not interaction_history or not indices_to_review:
212
+ raise ValueError("No unreviewed interactions to process.")
213
+ if verbose:
214
+ vizible.magenta(
215
+ f"Found {len(interaction_history)} unreviewed interactions."
216
+ )
217
+ interactions_to_review = []
218
+ for idx in indices_to_review:
219
+ # Mark as reviewed to skip re-processing in the next cycle.
220
+ interaction_history[idx].reviewed = True
221
+ interactions_to_review.append(interaction_history[idx])
222
+
223
+ # 1. Have the model re-evaluate its past responses and try to improve upon one of its turns.
224
+ review_prompt = revise.make_revision_prompt(
225
+ interactions_to_review, self._tokenizer
226
+ )
227
+ llm_rewrite_response = self.generate_response(review_prompt, use_history=False)
228
+ if verbose:
229
+ vizible.blue(f" - Response: {llm_rewrite_response}")
230
+
231
+ # 2. Prepare training data to train the model on how it should have responded in this
232
+ # situation.
233
+ example = revise.make_collated_training_example(
234
+ llm_rewrite_response, interactions_to_review, self._tokenizer
235
+ )
236
+ return example
237
+
238
+ def _train(self, example: TrainingExample, verbose: bool):
239
+ state = [self._model.state, self._optimizer.state, mx.random.state]
240
+ mx.eval(state)
241
+ loss_and_grad_fn = nn.value_and_grad(self._model, _loss_fn)
242
+
243
+ @functools.partial(mx.compile, inputs=state, outputs=state)
244
+ def _step(inputs, labels, mask):
245
+ loss, grads = loss_and_grad_fn(self._model, inputs, labels, mask)
246
+ self._optimizer.update(self._model, grads)
247
+ return loss
248
+
249
+ # 3. Train the model on the new, improved examples (backward-pass)
250
+ losses = []
251
+
252
+ if verbose:
253
+ vizible.green(f"During training: {mx.metal.device_info() = }")
254
+ mx.set_wired_limit(mx.metal.device_info()["max_recommended_working_set_size"])
255
+ mx.set_cache_limit(mx.metal.device_info()["max_buffer_length"])
256
+ world = mx.distributed.init()
257
+ world_size = world.size()
258
+ rank = world.rank()
259
+ if world_size > 1:
260
+ tqdm.tqdm.write(f"Node {rank} of {world_size}")
261
+
262
+ self._model.train(True)
263
+ for epoch in tqdm.tqdm(
264
+ range(self._epochs), desc="Training", total=self._epochs
265
+ ):
266
+ loss = _step(example.input, example.label, example.mask)
267
+ mx.eval(state, loss)
268
+ if verbose:
269
+ vizible.green(f"Epoch: {epoch}\tLoss: {loss = }")
270
+ losses.append(loss)
271
+
272
+ self._model.train(False)
273
+ if verbose:
274
+ vizible.cyan(f"{losses = }")
275
+
276
+ def self_correct_and_train(
277
+ self,
278
+ interaction_history: List[InteractionHistory],
279
+ indices_to_review: List[int] | None = None,
280
+ verbose: bool = False,
281
+ ) -> bool:
282
+ """Cycle in which the model revises a previously unreviewed prompt and trains from its rewrite.
283
+
284
+ Args:
285
+ interaction_history: Past user and model messages.
286
+ indices_to_review: Optional indices of relevant interactions to revise within
287
+ interaction_history. If not set, all interactions will be reviewed.
288
+ verbose: Enable verbose logging.
289
+
290
+ Returns:
291
+ Whether the process completed successfully.
292
+ """
293
+ self._model_is_stable = False
294
+
295
+ # Prepare training example from self-reflective revision of past dialog.
296
+ example = self._self_correct(interaction_history, indices_to_review, verbose)
297
+
298
+ # Train the model on the new, improved examples (backward-pass)
299
+ self._train(example, verbose)
300
+
301
+ self._model_is_stable = True
302
+ return True
@@ -0,0 +1,40 @@
1
+ """Server definition for spinning up a StatefulLLM instance."""
2
+
3
+ import asyncio
4
+ import logging
5
+ import socket
6
+ import sys
7
+
8
+ from fastapi import FastAPI
9
+ import uvicorn
10
+
11
+ from ._api import app
12
+
13
+ logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s', stream=sys.stdout)
14
+ logger = logging.getLogger(__name__) # Get a logger for this module
15
+
16
+
17
+ class MutableHostedLLM(uvicorn.Server):
18
+ """Host server to interact with a stateful LLM"""
19
+
20
+ def __init__(self, uvicorn_app: FastAPI = app, host: str = '127.0.0.1', port: int = 8000):
21
+ super().__init__(config=uvicorn.Config(uvicorn_app, host=host, port=port))
22
+ self._startup_done = asyncio.Event()
23
+ self._serve_task = None
24
+ self.should_exit = False
25
+
26
+ async def startup(self, sockets: list[socket.socket] | None = None) -> None:
27
+ """Override uvicorn startup"""
28
+ await super().startup(sockets=sockets)
29
+ self.config.setup_event_loop()
30
+ self._startup_done.set()
31
+
32
+ async def up(self) -> None:
33
+ """Start up server asynchronously"""
34
+ self._serve_task = asyncio.create_task(self.serve())
35
+ await self._startup_done.wait()
36
+
37
+ async def down(self) -> None:
38
+ """Shut down server asynchronously"""
39
+ self.should_exit = True
40
+ await self._serve_task
@@ -0,0 +1 @@
1
+ from . import revise