lora-easy 2.0__tar.gz → 2.2__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.
- {lora_easy-2.0 → lora_easy-2.2}/PKG-INFO +1 -1
- {lora_easy-2.0 → lora_easy-2.2}/lora_easy.egg-info/PKG-INFO +1 -1
- {lora_easy-2.0 → lora_easy-2.2}/lora_ez/agent.py +15 -6
- {lora_easy-2.0 → lora_easy-2.2}/lora_ez/model.py +93 -37
- {lora_easy-2.0 → lora_easy-2.2}/lora_ez/session.py +14 -6
- {lora_easy-2.0 → lora_easy-2.2}/pyproject.toml +1 -1
- {lora_easy-2.0 → lora_easy-2.2}/README.md +0 -0
- {lora_easy-2.0 → lora_easy-2.2}/lora_easy.egg-info/SOURCES.txt +0 -0
- {lora_easy-2.0 → lora_easy-2.2}/lora_easy.egg-info/dependency_links.txt +0 -0
- {lora_easy-2.0 → lora_easy-2.2}/lora_easy.egg-info/requires.txt +0 -0
- {lora_easy-2.0 → lora_easy-2.2}/lora_easy.egg-info/top_level.txt +0 -0
- {lora_easy-2.0 → lora_easy-2.2}/lora_ez/__init__.py +0 -0
- {lora_easy-2.0 → lora_easy-2.2}/lora_ez/commands.py +0 -0
- {lora_easy-2.0 → lora_easy-2.2}/lora_ez/deco.py +0 -0
- {lora_easy-2.0 → lora_easy-2.2}/setup.cfg +0 -0
|
@@ -29,8 +29,9 @@ class Agent:
|
|
|
29
29
|
name : str or None
|
|
30
30
|
Display name. Defaults to ``model.name``.
|
|
31
31
|
description : str or None
|
|
32
|
-
The agent's persona
|
|
33
|
-
|
|
32
|
+
The agent's persona. Defaults to ``model.description`` (the immutable
|
|
33
|
+
identity). Only forwarded as ``system_prompt`` when it differs from
|
|
34
|
+
the model's own description.
|
|
34
35
|
"""
|
|
35
36
|
|
|
36
37
|
def __init__(
|
|
@@ -51,7 +52,7 @@ class Agent:
|
|
|
51
52
|
):
|
|
52
53
|
self.model = model
|
|
53
54
|
self.name = name or model.name
|
|
54
|
-
self.description = description or model.
|
|
55
|
+
self.description = description or model.description
|
|
55
56
|
|
|
56
57
|
self.web_enabled = web_enabled
|
|
57
58
|
self.web_allowlist = web_allowlist or []
|
|
@@ -81,7 +82,7 @@ class Agent:
|
|
|
81
82
|
return cls(
|
|
82
83
|
model,
|
|
83
84
|
name=cfg.get("name", model.name),
|
|
84
|
-
description=cfg.get("description", model.
|
|
85
|
+
description=cfg.get("description", model.description),
|
|
85
86
|
web_enabled=cfg.get("web_enabled", False),
|
|
86
87
|
web_allowlist=cfg.get("web_allowlist"),
|
|
87
88
|
web_blocklist=cfg.get("web_blocklist"),
|
|
@@ -115,9 +116,14 @@ class Agent:
|
|
|
115
116
|
f"{prompt}\n\n"
|
|
116
117
|
f"[Relevant context from tools:\n{context}\n]"
|
|
117
118
|
)
|
|
119
|
+
# model.description is prepended by the model itself; only forward a
|
|
120
|
+
# system_prompt that differs from it (avoids duplicating the identity)
|
|
121
|
+
sp = system_prompt
|
|
122
|
+
if sp is None and self.description and self.description != self.model.description:
|
|
123
|
+
sp = self.description
|
|
118
124
|
return self.model.chat(
|
|
119
125
|
prompt, history=history, max_tokens=max_tokens,
|
|
120
|
-
system_prompt=
|
|
126
|
+
system_prompt=sp, **kwargs,
|
|
121
127
|
)
|
|
122
128
|
|
|
123
129
|
# -- context gathering ---------------------------------------------------
|
|
@@ -223,8 +229,11 @@ class Agent:
|
|
|
223
229
|
|
|
224
230
|
def chat_session(self, *args, **kwargs):
|
|
225
231
|
from .session import _ChatSession
|
|
232
|
+
sp = kwargs.pop("system_prompt", None)
|
|
233
|
+
if sp is None and self.description != self.model.description:
|
|
234
|
+
sp = self.description
|
|
226
235
|
return _ChatSession(
|
|
227
236
|
self.model, *args,
|
|
228
|
-
system_prompt=
|
|
237
|
+
system_prompt=sp,
|
|
229
238
|
**kwargs,
|
|
230
239
|
)
|
|
@@ -22,9 +22,8 @@ Use::
|
|
|
22
22
|
s.run() # interactive REPL, /exit to quit
|
|
23
23
|
"""
|
|
24
24
|
|
|
25
|
-
from pathlib import Path
|
|
26
|
-
|
|
27
25
|
import torch
|
|
26
|
+
from pathlib import Path
|
|
28
27
|
from peft import LoraConfig, get_peft_model, PeftModel
|
|
29
28
|
from transformers import AutoModelForCausalLM, AutoTokenizer, Trainer, TrainingArguments
|
|
30
29
|
|
|
@@ -34,7 +33,7 @@ class LoraModel:
|
|
|
34
33
|
|
|
35
34
|
Parameters
|
|
36
35
|
----------
|
|
37
|
-
|
|
36
|
+
id_ : str
|
|
38
37
|
HuggingFace model identifier (e.g. ``"Qwen/Qwen2.5-0.5B-Instruct"``).
|
|
39
38
|
name : str
|
|
40
39
|
Human-readable label, used for default save paths and REPL display.
|
|
@@ -43,6 +42,16 @@ class LoraModel:
|
|
|
43
42
|
system_prompt : str or None
|
|
44
43
|
Prepend a system instruction to every chat turn. When ``None`` the
|
|
45
44
|
model uses its built-in default (``"You are Qwen …"``).
|
|
45
|
+
description : str or None
|
|
46
|
+
Immutable identity/role (e.g. ``"you are a sassy house cat"``).
|
|
47
|
+
Rendered as its own system message BEFORE ``system_prompt``.
|
|
48
|
+
Saved with the adapter (``description.txt``) and restored on ``load()``.
|
|
49
|
+
``None`` (default) means no description is stored or loaded.
|
|
50
|
+
save_path : str
|
|
51
|
+
The path where the adapter is saved.
|
|
52
|
+
cache_dir : str or None
|
|
53
|
+
Custom model download cache. ``None`` uses the default
|
|
54
|
+
``~/.cache/huggingface/hub/``.
|
|
46
55
|
|
|
47
56
|
Key attributes
|
|
48
57
|
--------------
|
|
@@ -54,34 +63,46 @@ class LoraModel:
|
|
|
54
63
|
The active model — ``base`` or ``peft_model`` depending on LoRA state.
|
|
55
64
|
"""
|
|
56
65
|
|
|
57
|
-
def __init__(self,
|
|
58
|
-
system_prompt: str | None = None
|
|
59
|
-
|
|
66
|
+
def __init__(self, id_: str, name: str = "Assistant", device: str = "mps",
|
|
67
|
+
system_prompt: str | None = None, save_path: str | None = None,
|
|
68
|
+
cache_dir: str | None = None, description: str | None = None):
|
|
69
|
+
self.id_ = id_
|
|
60
70
|
self.name = name
|
|
61
71
|
self.system_prompt = system_prompt
|
|
62
|
-
self.
|
|
72
|
+
self._description = description
|
|
73
|
+
# tokenizer download (cached at cache_dir or ~/.cache/huggingface/hub/)
|
|
74
|
+
self.tokenizer = AutoTokenizer.from_pretrained(id_, cache_dir=cache_dir)
|
|
63
75
|
self.tokenizer.pad_token = self.tokenizer.eos_token
|
|
64
|
-
|
|
76
|
+
# model weights download (largest; first run downloads GBs)
|
|
77
|
+
self.base = AutoModelForCausalLM.from_pretrained(id_, device_map=device, torch_dtype="auto", cache_dir=cache_dir)
|
|
65
78
|
self.peft_model = None
|
|
66
79
|
self.model = self.base
|
|
67
80
|
self._lora_enabled = False
|
|
81
|
+
self.save_path = save_path or f"{self:l}-lora"
|
|
82
|
+
|
|
83
|
+
@property
|
|
84
|
+
def description(self) -> str | None:
|
|
85
|
+
"""Immutable role/identity string (read-only; set via ``__init__`` or ``load()``)."""
|
|
86
|
+
return self._description
|
|
68
87
|
|
|
69
88
|
def __repr__(self) -> str:
|
|
70
89
|
n = sum(p.numel() for p in self.model.parameters())
|
|
71
|
-
return f"{self
|
|
90
|
+
return f"{self:t}: LoraModel('{self:i}', {n/1e9:.1f}B params, lora={'yes' if self._has_lora() else 'no'})"
|
|
72
91
|
|
|
73
92
|
def __str__(self):
|
|
74
93
|
return self.name.title()
|
|
75
94
|
|
|
76
|
-
def __format__(self, spec=
|
|
77
|
-
if spec
|
|
95
|
+
def __format__(self, spec=''):
|
|
96
|
+
if spec == '':
|
|
78
97
|
return self.name
|
|
79
|
-
elif spec
|
|
98
|
+
elif spec == 'l':
|
|
80
99
|
return self.name.lower()
|
|
81
|
-
elif spec
|
|
100
|
+
elif spec == 't':
|
|
82
101
|
return self.name
|
|
102
|
+
elif spec == 'i':
|
|
103
|
+
return self.id_
|
|
83
104
|
else:
|
|
84
|
-
raise ValueError('`spec` should be one of
|
|
105
|
+
raise ValueError('`spec` should be one of [empty] | l | t')
|
|
85
106
|
|
|
86
107
|
def _has_lora(self) -> bool:
|
|
87
108
|
return self.peft_model is not None
|
|
@@ -130,12 +151,19 @@ class LoraModel:
|
|
|
130
151
|
``history`` is a list of ``{"role": ..., "content": ...}`` dicts from
|
|
131
152
|
previous turns. The model sees the full context and can refer back to it.
|
|
132
153
|
``system_prompt`` overrides ``self.system_prompt`` for this single turn.
|
|
154
|
+
``description`` (if set) is always prepended as its own system message.
|
|
133
155
|
"""
|
|
134
156
|
self.model.eval()
|
|
135
157
|
sp = system_prompt if system_prompt is not None else self.system_prompt
|
|
136
158
|
messages = list(history or [])
|
|
137
|
-
if
|
|
138
|
-
|
|
159
|
+
if not messages or messages[0].get("role") != "system":
|
|
160
|
+
# description first (immutable identity), then system_prompt (mutable)
|
|
161
|
+
sys_msgs = []
|
|
162
|
+
if self.description:
|
|
163
|
+
sys_msgs.append({"role": "system", "content": self.description})
|
|
164
|
+
if sp:
|
|
165
|
+
sys_msgs.append({"role": "system", "content": sp})
|
|
166
|
+
messages = sys_msgs + messages
|
|
139
167
|
messages.append({"role": "user", "content": prompt})
|
|
140
168
|
fmt = self._format(messages)
|
|
141
169
|
inp = self.tokenizer(fmt, return_tensors="pt").to(self.model.device)
|
|
@@ -157,13 +185,12 @@ class LoraModel:
|
|
|
157
185
|
Usage::
|
|
158
186
|
|
|
159
187
|
with model.chat_session("./history.json", auto_save=True) as s:
|
|
160
|
-
s
|
|
161
|
-
s
|
|
188
|
+
s> "Hello"
|
|
189
|
+
s> "What do you think?"
|
|
162
190
|
s.history # all turns so far
|
|
163
191
|
|
|
164
192
|
When the history grows beyond ``max_tokens`` tokens it is automatically
|
|
165
193
|
summarised by the model and replaced with a compact system message.
|
|
166
|
-
``system_prompt`` overrides ``self.system_prompt`` for the session.
|
|
167
194
|
"""
|
|
168
195
|
from .session import _ChatSession
|
|
169
196
|
return _ChatSession(self, save_path, auto_save, max_tokens,
|
|
@@ -171,7 +198,8 @@ class LoraModel:
|
|
|
171
198
|
|
|
172
199
|
# -- Training -----------------------------------------------------------
|
|
173
200
|
|
|
174
|
-
def train(self, conversations: list[dict],
|
|
201
|
+
def train(self, conversations: list[dict], save_checkpoints: bool = False,
|
|
202
|
+
max_length: int = 256, **kwargs):
|
|
175
203
|
"""Fine-tune with LoRA on ShareGPT-format conversations.
|
|
176
204
|
|
|
177
205
|
Pipeline: render each convo via chat template -> tokenize with pad+trunc ->
|
|
@@ -179,30 +207,39 @@ class LoraModel:
|
|
|
179
207
|
Auto-enables LoRA if not already active.
|
|
180
208
|
|
|
181
209
|
Args:
|
|
182
|
-
|
|
183
|
-
False (default) produces no
|
|
210
|
+
save_checkpoints: if set, save a checkpoint per epoch to
|
|
211
|
+
``./lora-output-{name}``. False (default) produces no files.
|
|
212
|
+
max_length: tokenizer truncation/padding length. Raise it when
|
|
213
|
+
conversations are long (multi-turn or long replies) —
|
|
214
|
+
truncation cuts from the END, which would clip the target
|
|
215
|
+
assistant reply. Default 256.
|
|
184
216
|
"""
|
|
185
217
|
if not self.lora_enabled:
|
|
186
218
|
self.enable_lora()
|
|
187
219
|
self.model.config.use_cache = False
|
|
188
220
|
|
|
189
221
|
texts = [self._format(self._render(c), add_gen=False) for c in conversations]
|
|
190
|
-
tok = self.tokenizer(texts, truncation=True, padding="max_length",
|
|
222
|
+
tok = self.tokenizer(texts, truncation=True, padding="max_length",
|
|
223
|
+
max_length=max_length)
|
|
191
224
|
dataset = [{"input_ids": inds, "attention_mask": ms,
|
|
192
225
|
"labels": [-100 if m == 0 else i for i, m in zip(inds, ms)]}
|
|
193
226
|
for inds, ms in zip(tok["input_ids"], tok["attention_mask"])]
|
|
194
227
|
|
|
195
|
-
|
|
196
|
-
defaults = {"
|
|
197
|
-
|
|
228
|
+
# friendly defaults; translate to TrainingArguments names below
|
|
229
|
+
defaults = {"output_dir": f"./lora-output-{self:l}" if save_checkpoints else "./temp_output",
|
|
230
|
+
"epochs": 30, "lr": 3e-4,
|
|
231
|
+
"per_device_train_batch_size": 4, "logging_steps": 5}
|
|
232
|
+
kwargs = defaults | kwargs
|
|
198
233
|
args = TrainingArguments(
|
|
199
|
-
|
|
200
|
-
|
|
234
|
+
num_train_epochs=kwargs.pop("epochs"),
|
|
235
|
+
learning_rate=kwargs.pop("lr"),
|
|
236
|
+
**kwargs,
|
|
237
|
+
save_strategy="no" if not save_checkpoints else "epoch",
|
|
201
238
|
report_to="none",
|
|
202
239
|
)
|
|
203
240
|
Trainer(model=self.model, args=args, train_dataset=dataset, processing_class=self.tokenizer).train()
|
|
204
241
|
|
|
205
|
-
if not
|
|
242
|
+
if not save_checkpoints:
|
|
206
243
|
import shutil
|
|
207
244
|
shutil.rmtree("./temp_output", ignore_errors=True)
|
|
208
245
|
|
|
@@ -218,24 +255,43 @@ class LoraModel:
|
|
|
218
255
|
def from_yaml(cls, path: str):
|
|
219
256
|
"""Create a LoraModel from a YAML config file.
|
|
220
257
|
|
|
221
|
-
Expected keys: ``
|
|
258
|
+
Expected keys: ``id_``, ``name``, ``system_prompt``,
|
|
222
259
|
``peft_path`` (optional — auto-loads the adapter).
|
|
223
260
|
"""
|
|
224
261
|
import yaml
|
|
262
|
+
from pathlib import Path
|
|
225
263
|
cfg = yaml.safe_load(Path(path).read_text())
|
|
226
|
-
|
|
227
|
-
|
|
264
|
+
if "id_" in cfg:
|
|
265
|
+
id_ = cfg["id_"]
|
|
266
|
+
elif "id" in cfg:
|
|
267
|
+
id_ = cfg["id"]
|
|
268
|
+
else:
|
|
269
|
+
raise KeyError("Not provide the key `id_` | `id`.")
|
|
270
|
+
m = cls(id_=id_, name=cfg.get("name", "Assistant"),
|
|
271
|
+
system_prompt=cfg.get("system_prompt"),
|
|
272
|
+
description=cfg.get("description"))
|
|
228
273
|
if peft := cfg.get("peft_path"):
|
|
229
274
|
m.load(peft)
|
|
230
275
|
return m
|
|
231
276
|
|
|
232
277
|
def save(self, path: str | None = None):
|
|
233
|
-
|
|
278
|
+
# save the adapter
|
|
279
|
+
path = path or self.save_path
|
|
234
280
|
self.model.save_pretrained(path)
|
|
235
281
|
self.tokenizer.save_pretrained(path)
|
|
236
|
-
|
|
282
|
+
if self.description is not None:
|
|
283
|
+
(Path(path) / "description.txt").write_text(self.description)
|
|
237
284
|
|
|
238
285
|
def load(self, path: str | None = None):
|
|
239
|
-
|
|
240
|
-
|
|
241
|
-
|
|
286
|
+
# load the adapter from a LOCAL directory
|
|
287
|
+
p = Path(path or self.save_path)
|
|
288
|
+
if not p.is_dir():
|
|
289
|
+
raise FileNotFoundError(
|
|
290
|
+
f"""Adapter directory not found: '{p}'.
|
|
291
|
+
If it lives elsewhere pass the path explicitly, e.g. m.load('/path/to/ivanka-lora').""")
|
|
292
|
+
self.peft_model = PeftModel.from_pretrained(self.base, str(p),
|
|
293
|
+
is_trainable=True)
|
|
294
|
+
# restore description if it was saved alongside the adapter
|
|
295
|
+
desc_file = p / "description.txt"
|
|
296
|
+
if desc_file.exists():
|
|
297
|
+
self._description = desc_file.read_text()
|
|
@@ -73,11 +73,15 @@ class _ChatSession:
|
|
|
73
73
|
"""
|
|
74
74
|
self._maybe_summarize(prompt)
|
|
75
75
|
self.history.append({"role": "user", "content": prompt})
|
|
76
|
-
reply = self.model.chat(prompt, history=self.history[:-1],
|
|
76
|
+
reply = self.model.chat(prompt, history=self.history[:-1],
|
|
77
77
|
system_prompt=self.system_prompt, **kwargs)
|
|
78
78
|
self.history.append({"role": "assistant", "content": reply})
|
|
79
79
|
return reply
|
|
80
80
|
|
|
81
|
+
def __gt__(self, prompt: str) -> str:
|
|
82
|
+
"""``s > "prompt"`` is shorthand for ``s.chat("prompt")``."""
|
|
83
|
+
return self.chat(prompt)
|
|
84
|
+
|
|
81
85
|
def run(self):
|
|
82
86
|
"""Interactive REPL — type messages at a prompt, ``/exit`` to quit."""
|
|
83
87
|
print(f"Chat session started. Type /exit to quit.\n"
|
|
@@ -127,10 +131,14 @@ class _ChatSession:
|
|
|
127
131
|
raw = self.model.tokenizer.decode(out[0], skip_special_tokens=True)
|
|
128
132
|
summary = raw.rpartition("assistant\n")[-1].strip()
|
|
129
133
|
|
|
130
|
-
# preserve
|
|
134
|
+
# preserve description (immutable) as its own system message,
|
|
135
|
+
# then merge summary into the system_prompt (mutable)
|
|
136
|
+
self.history = []
|
|
137
|
+
if self.model.description:
|
|
138
|
+
self.history.append({"role": "system", "content": self.model.description})
|
|
131
139
|
if self.system_prompt:
|
|
132
|
-
self.history
|
|
133
|
-
"content": f"{self.system_prompt}\n\n[Summary of previous turns: {summary}]"}
|
|
140
|
+
self.history.append({"role": "system",
|
|
141
|
+
"content": f"{self.system_prompt}\n\n[Summary of previous turns: {summary}]"})
|
|
134
142
|
else:
|
|
135
|
-
self.history
|
|
136
|
-
"content": f"Previous conversation summary: {summary}"}
|
|
143
|
+
self.history.append({"role": "system",
|
|
144
|
+
"content": f"Previous conversation summary: {summary}"})
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|