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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: lora-easy
3
- Version: 2.0
3
+ Version: 2.2
4
4
  Summary: A tiny OOP wrapper around PEFT for LoRA fine-tuning of causal LMs.
5
5
  Author: Freakwill
6
6
  License: MIT
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: lora-easy
3
- Version: 2.0
3
+ Version: 2.2
4
4
  Summary: A tiny OOP wrapper around PEFT for LoRA fine-tuning of causal LMs.
5
5
  Author: Freakwill
6
6
  License: MIT
@@ -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 — forwarded as ``system_prompt`` on every turn.
33
- Defaults to ``model.system_prompt``.
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.system_prompt
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.system_prompt),
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=system_prompt or self.description, **kwargs,
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=kwargs.pop("system_prompt", self.description),
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
- model_id : str
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, model_id: str, name: str = "Assistant", device: str = "mps",
58
- system_prompt: str | None = None):
59
- self.model_id = model_id
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.tokenizer = AutoTokenizer.from_pretrained(model_id)
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
- self.base = AutoModelForCausalLM.from_pretrained(model_id, device_map=device, torch_dtype="auto")
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.name.title()}: LoraModel('{self.model_id}', {n/1e9:.1f}B params, lora={'yes' if self._has_lora() else 'no'})"
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=None):
77
- if spec is None:
95
+ def __format__(self, spec=''):
96
+ if spec == '':
78
97
  return self.name
79
- elif spec is 'l':
98
+ elif spec == 'l':
80
99
  return self.name.lower()
81
- elif spec is 't':
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 None | l | t')
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 sp and (not messages or messages[0].get("role") != "system"):
138
- messages = [{"role": "system", "content": sp}] + messages
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.chat("Hello")
161
- s.chat("What do you think?")
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], output: bool = False, **kwargs):
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
- output: if set, save training logs/checkpoints to ``./lora-output-{name}``.
183
- False (default) produces no output files.
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", max_length=256)
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
- output_dir = f"./lora-output-{self.name.lower()}" if output else "./temp_output"
196
- defaults = {"epochs": 30, "lr": 3e-4, "per_device_train_batch_size": 4, "logging_steps": 5}
197
- merged = defaults | kwargs
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
- output_dir=output_dir, **merged,
200
- save_strategy="no" if output is None else "epoch",
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 output:
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: ``model_id``, ``name``, ``system_prompt``,
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
- m = cls(cfg["model_id"], name=cfg.get("name", "Assistant"),
227
- system_prompt=cfg.get("system_prompt"))
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
- path = f"lora-{self:l}" if path is None else path
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
- print(f"The adapter is saved in `{path}`.")
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
- path = f"lora-{self:l}" if path is None else path
240
- self.peft_model = PeftModel.from_pretrained(self.base, path)
241
- # self.enable_lora()
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 system prompt if set, merge summary into it
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 = [{"role": "system",
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 = [{"role": "system",
136
- "content": f"Previous conversation summary: {summary}"}]
143
+ self.history.append({"role": "system",
144
+ "content": f"Previous conversation summary: {summary}"})
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "lora-easy"
7
- version = "2.0"
7
+ version = "2.2"
8
8
  description = "A tiny OOP wrapper around PEFT for LoRA fine-tuning of causal LMs."
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.9"
File without changes
File without changes
File without changes
File without changes
File without changes