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.
- adaptible-1.0.0a1/.gitignore +4 -0
- adaptible-1.0.0a1/PKG-INFO +17 -0
- adaptible-1.0.0a1/README.md +0 -0
- adaptible-1.0.0a1/adaptible/__init__.py +3 -0
- adaptible-1.0.0a1/adaptible/_src/__init__.py +6 -0
- adaptible-1.0.0a1/adaptible/_src/_api.py +142 -0
- adaptible-1.0.0a1/adaptible/_src/_classes.py +80 -0
- adaptible-1.0.0a1/adaptible/_src/_llm.py +302 -0
- adaptible-1.0.0a1/adaptible/_src/_server.py +40 -0
- adaptible-1.0.0a1/adaptible/_src/libs/__init__.py +1 -0
- adaptible-1.0.0a1/adaptible/_src/libs/revise.py +211 -0
- adaptible-1.0.0a1/adaptible/_src/static/favicon.ico +0 -0
- adaptible-1.0.0a1/adaptible/_src/static/home.html +381 -0
- adaptible-1.0.0a1/adaptible/_src/static/ible.png +0 -0
- adaptible-1.0.0a1/adaptible/_src/static/index.html +254 -0
- adaptible-1.0.0a1/examples/online_learning_demo.py +107 -0
- adaptible-1.0.0a1/examples/server_client.ipynb +172 -0
- adaptible-1.0.0a1/examples/server_demo.py +31 -0
- adaptible-1.0.0a1/examples/web_app_demo.py +44 -0
- adaptible-1.0.0a1/pyproject.toml +30 -0
- adaptible-1.0.0a1/tests/llm_test.py +48 -0
|
@@ -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,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
|