RegisterAdapterInOllama(string baseUrl, string outputName, string ollamaBase, string adapterGguf)
+ {
StringBuilder mf = new();
- JObject job = Memory.GetTrainJob(TrainingJobManager.CurrentJobId ?? "");
- string baseModel = job?["base_model"]?.ToString() ?? "unknown";
- mf.AppendLine($"FROM {baseModel}");
- mf.AppendLine($"ADAPTER {adapterFile.Replace("\\", "/")}");
+ mf.AppendLine($"FROM {ollamaBase}");
+ mf.AppendLine($"ADAPTER {adapterGguf.Replace("\\", "/")}");
JObject payload = new()
{
["name"] = outputName,
@@ -224,14 +373,31 @@ public partial class SwarmAssistentExtension
};
using StringContent content = new(payload.ToString(Newtonsoft.Json.Formatting.None), Encoding.UTF8, "application/json");
using HttpResponseMessage resp = await HttpClient.PostAsync($"{NormalizeBaseUrl(baseUrl)}/api/create", content);
- _ = await resp.Content.ReadAsStringAsync();
+ string body = await resp.Content.ReadAsStringAsync();
+ if (!resp.IsSuccessStatusCode)
+ {
+ return new JObject
+ {
+ ["success"] = false,
+ ["error"] = $"ollama create HTTP {(int)resp.StatusCode}: {Clip(body, 400)}",
+ ["modelfile"] = mf.ToString(),
+ };
+ }
+ return new JObject
+ {
+ ["success"] = true,
+ ["name"] = outputName,
+ ["ollama_base"] = ollamaBase,
+ ["adapter"] = adapterGguf,
+ ["response"] = body,
+ };
}
}
sealed class TrainingJobManager
{
static readonly Regex LossRe = new(@"loss[:\s]+([0-9.]+)", RegexOptions.IgnoreCase | RegexOptions.Compiled);
- static readonly Regex StepRe = new(@"(\d+)\s*/\s*(\d+)", RegexOptions.Compiled);
+ static readonly Regex StepRe = new(@"step\s+(\d+)\s*/\s*(\d+)", RegexOptions.IgnoreCase | RegexOptions.Compiled);
Process _process;
readonly object _lock = new();
@@ -242,6 +408,8 @@ sealed class TrainingJobManager
string _jobId;
string _baseUrl;
string _chatModel;
+ long _lastProgressSaveMs;
+ int _lastSavedStep = -1;
public bool IsRunning { get; private set; }
public string CurrentJobId => _jobId;
@@ -260,6 +428,8 @@ sealed class TrainingJobManager
_logPath = logPath;
_baseUrl = baseUrl;
_chatModel = chatModel;
+ _lastProgressSaveMs = 0;
+ _lastSavedStep = -1;
_progress = new JObject { ["status"] = "running", ["step"] = 0, ["loss"] = null, ["log"] = "" };
try
{
@@ -276,6 +446,7 @@ sealed class TrainingJobManager
if (!string.IsNullOrWhiteSpace(hfToken))
{
psi.Environment["HF_TOKEN"] = hfToken;
+ psi.Environment["HUGGING_FACE_HUB_TOKEN"] = hfToken;
}
_process = new Process { StartInfo = psi, EnableRaisingEvents = true };
_process.OutputDataReceived += (_, e) => AppendLog(e.Data);
@@ -313,7 +484,7 @@ sealed class TrainingJobManager
// ignore
}
string prev = _progress["log"]?.ToString() ?? "";
- string combined = (prev + line + "\n");
+ string combined = prev + line + "\n";
if (combined.Length > 12000)
{
combined = combined[^12000..];
@@ -327,33 +498,48 @@ sealed class TrainingJobManager
Match stepM = StepRe.Match(line);
if (stepM.Success)
{
- _progress["step"] = int.Parse(stepM.Groups[1].Value);
- _progress["total_steps"] = int.Parse(stepM.Groups[2].Value);
- int total = int.Parse(stepM.Groups[2].Value);
int step = int.Parse(stepM.Groups[1].Value);
+ int total = int.Parse(stepM.Groups[2].Value);
+ _progress["step"] = step;
+ _progress["total_steps"] = total;
_progress["percent"] = total > 0 ? (int)(100.0 * step / total) : 0;
}
- try
+ MaybeSaveProgress(stepM.Success ? int.Parse(stepM.Groups[1].Value) : -1);
+ }
+ }
+
+ void MaybeSaveProgress(int step)
+ {
+ long now = DateTimeOffset.UtcNow.ToUnixTimeMilliseconds();
+ bool stepChanged = step >= 0 && step != _lastSavedStep;
+ if (!stepChanged && now - _lastProgressSaveMs < 2500)
+ {
+ return;
+ }
+ _lastProgressSaveMs = now;
+ if (step >= 0)
+ {
+ _lastSavedStep = step;
+ }
+ try
+ {
+ _ext?.Memory?.SaveTrainJob(new JObject
{
- _ext?.Memory?.SaveTrainJob(new JObject
- {
- ["id"] = _jobId,
- ["status"] = "running",
- ["progress"] = _progress,
- });
- }
- catch
- {
- // ignore
- }
+ ["id"] = _jobId,
+ ["status"] = "running",
+ ["progress"] = _progress,
+ });
+ }
+ catch
+ {
+ // ignore
}
}
async Task OnExited()
{
bool ok = false;
- string adapterDir = "";
- string outputName = "";
+ JObject jobConfig = new();
lock (_lock)
{
ok = _process?.ExitCode == 0;
@@ -366,16 +552,18 @@ sealed class TrainingJobManager
JObject job = _ext.Memory.GetTrainJob(_jobId);
try
{
- JObject cfg = JObject.Parse(job?["config_json"]?.ToString() ?? "{}");
- adapterDir = cfg["adapter_dir"]?.ToString() ?? "";
- outputName = cfg["output_name"]?.ToString() ?? job?["output_name"]?.ToString() ?? "";
+ string cfgRaw = job?["config_json"]?.ToString();
+ if (!string.IsNullOrWhiteSpace(cfgRaw))
+ {
+ jobConfig = JObject.Parse(cfgRaw);
+ }
}
catch
{
// ignore
}
JObject runner = _ext.Config.LoadTrainingRunner();
- await _ext.FinishTrainJobAsync(_jobId, ok, _logPath, _session, _baseUrl, _chatModel, adapterDir, outputName, runner["gguf_script"]?.ToString());
+ await _ext.FinishTrainJobAsync(_jobId, ok, _logPath, _session, _baseUrl, _chatModel, jobConfig, runner);
}
}
diff --git a/README.md b/README.md
index 9a1e26a..7253b98 100644
--- a/README.md
+++ b/README.md
@@ -4,6 +4,8 @@ SwarmUI extension for **collaborative Krea 2** prompting via **Ollama**: chat +
**Turn model:** one user message is one *turn*. A turn may fan out into nested LLM *hops* — Krea prompt prep, empty-patch retry, vision, auto-critique. Hops share one `HOP_BUDGET`, never re-read the user's text (their prompt is client-authored), and pass the busy gate that blocks new user sends. What a reply does to generation state is decided once, in `resolveTurnIntent`: the model's `actions:["generate"]` / `look_at` win, RU intent heuristics only back it up when the model forgets, and an explicit «запомни, не генерируй» vetoes both.
+**Version 0.13.0** — **Реальный QLoRA-пайплайн**: `train_qlora.py` (TRL SFTTrainer + PEFT), HF-датасеты с маппингом (preset fiction title/tags→text), `max_samples`, полный post-train: safetensors → GGUF (`convert_lora_to_gguf.py`) → `ollama create` с `FROM ollama_base` + `ADAPTER`. Раннер: `builtin` + `custom`. Зависимости: `scripts/requirements-train.txt`.
+
**Version 0.12.1** — **Услышанное → агент**: одобренные примеры датасета сразу попадают в vector memory (`kind=heard`) и в контекст чата как `heard_examples` (без QLoRA). На вкладке «Датасет»: авто-подключение при одобрении, синхронизация всех, per-sample 🔗. Агент может запросить `heard_search`. Настройки: `training-agent.json`.
**Version 0.12.0** — App-level tabs (Чат / Карточки / **Обучение** / Настройки), боковая панель истории чатов, вкладка обучения LLM: курирование диалогов, импорт JSONL/CSV, Hugging Face datasets (фильтр совместимости), быстрый Ollama Modelfile, опциональный QLoRA-раннер с локаутом VRAM. HF token из SwarmUI User Settings (`huggingface_api`).
@@ -217,6 +219,16 @@ Patch fence keys: single source `Config/_base/patch-keys.json` → C# + client v
| `AssistentGetDatasetAgentSettings` / `AssistentSaveDatasetAgentSettings` | «Услышанное» → agent RAG (`training-agent.json`) |
| `AssistentLinkTrainSampleToAgent` / `AssistentUnlinkTrainSampleFromAgent` / `AssistentSyncDatasetToAgent` | Embed approved samples as `heard` memory |
+## QLoRA setup (0.13.0)
+
+1. Python env with CUDA: `pip install -r scripts/requirements-train.txt`
+2. SwarmUI **User Settings** → `huggingface_api` (for HF base model download)
+3. **Настройки → Модели**: runner kind = **builtin**, paths to `convert_lora_to_gguf.py` and **GGUF base** (same arch as HF base)
+4. **Обучение → QLoRA**: HF base id, **Ollama base** (existing tag), output name, optional HF dataset (`krplt/ru-fictext-nsfw` auto-maps fiction preset)
+5. Pipeline: train → `adapter_model.safetensors` → GGUF → `ollama create` with `ADAPTER`
+
+Manual test checklist: small JSONL (5 pairs); HF dataset with `max_samples=50`; cancel job; missing deps (exit 2); missing gguf script (completed with note).
+
## License
MIT
diff --git a/SwarmAssistentExtension.cs b/SwarmAssistentExtension.cs
index e7154ab..d1d22d9 100644
--- a/SwarmAssistentExtension.cs
+++ b/SwarmAssistentExtension.cs
@@ -33,8 +33,8 @@ public partial class SwarmAssistentExtension : Extension
ExtensionAuthor = "mrleo1nid";
Description = "Collaborative Krea 2 assistant: Ollama chat, persona presets, vector memory, model cards, Generate loop.";
License = "MIT";
- Version = "0.12.1";
- Tags = ["tabs", "ui", "llm", "ollama", "krea", "inpaint", "memory", "training", "heard"];
+ Version = "0.13.0";
+ Tags = ["tabs", "ui", "llm", "ollama", "krea", "inpaint", "memory", "training", "heard", "qlora"];
}
public override void OnInit()
@@ -104,7 +104,7 @@ public partial class SwarmAssistentExtension : Extension
API.RegisterAPICall(AssistentLinkTrainSampleToAgent, true, PermUse);
API.RegisterAPICall(AssistentUnlinkTrainSampleFromAgent, true, PermUse);
API.RegisterAPICall(AssistentSyncDatasetToAgent, true, PermUse);
- Logs.Init("Swarm Assistent extension loaded (0.12.1 heard dataset → agent)");
+ Logs.Init("Swarm Assistent extension loaded (0.13.0 real QLoRA pipeline)");
}
int CfgInt(string key, int fallback)
diff --git a/Tabs/Text2Image/Assistent.html b/Tabs/Text2Image/Assistent.html
index 221fe25..72d48da 100644
--- a/Tabs/Text2Image/Assistent.html
+++ b/Tabs/Text2Image/Assistent.html
@@ -222,6 +222,17 @@
+
+ Маппинг колонок:
+
+
+
+
@@ -253,17 +264,22 @@
@@ -323,15 +339,15 @@
+
-
+
+
diff --git a/scripts/hf_dataset_map.py b/scripts/hf_dataset_map.py
new file mode 100644
index 0000000..5d21be7
--- /dev/null
+++ b/scripts/hf_dataset_map.py
@@ -0,0 +1,94 @@
+"""Map Hugging Face dataset rows to chat messages for SFT."""
+from __future__ import annotations
+
+
+def fiction_tags_text_user(row: dict, title_col: str = "title", tags_col: str = "tags") -> str:
+ parts = []
+ title = (row.get(title_col) or "").strip()
+ tags = (row.get(tags_col) or "").strip()
+ if title:
+ parts.append(f"Title: {title}")
+ if tags:
+ parts.append(f"Tags: {tags}")
+ return "\n".join(parts) if parts else tags or title or ""
+
+
+def row_to_messages(row: dict, schema: dict | None, mapping: dict | None) -> list[dict] | None:
+ schema = schema or {}
+ mapping = mapping or {}
+ kind = schema.get("kind") or mapping.get("kind") or mapping.get("preset")
+
+ if kind == "messages" and row.get("messages"):
+ return _normalize_messages(row["messages"])
+
+ if kind == "conversations" and row.get("conversations"):
+ out = []
+ for c in row["conversations"]:
+ if not isinstance(c, dict):
+ continue
+ frm = c.get("from") or ""
+ val = c.get("value") or ""
+ role = "assistant" if frm in ("gpt", "assistant", "chatgpt") else "user"
+ if frm in ("human", "user"):
+ role = "user"
+ if val:
+ out.append({"role": role, "content": str(val)})
+ return out or None
+
+ if kind == "alpaca":
+ instr = (row.get("instruction") or "").strip()
+ inp = (row.get("input") or "").strip()
+ output = (row.get("output") or "").strip()
+ user = instr if not inp else f"{instr}\n{inp}"
+ if user and output:
+ return [{"role": "user", "content": user}, {"role": "assistant", "content": output}]
+ return None
+
+ if kind == "prompt_response":
+ resp_col = schema.get("response_col") or "response"
+ prompt = row.get("prompt") or ""
+ resp = row.get(resp_col) or ""
+ if prompt and resp:
+ return [{"role": "user", "content": str(prompt)}, {"role": "assistant", "content": str(resp)}]
+ return None
+
+ if kind == "qa":
+ q = row.get("question") or ""
+ a = row.get("answer") or ""
+ if q and a:
+ return [{"role": "user", "content": str(q)}, {"role": "assistant", "content": str(a)}]
+ return None
+
+ if kind in ("fiction_tags_text", "preset_fiction_tags_text"):
+ text = (row.get("text") or "").strip()
+ user = fiction_tags_text_user(row)
+ if user and text:
+ return [{"role": "user", "content": user}, {"role": "assistant", "content": text}]
+ return None
+
+ if kind == "custom" or mapping.get("user_col"):
+ user_col = mapping.get("user_col")
+ asst_col = mapping.get("assistant_col")
+ if user_col and asst_col:
+ u = row.get(user_col) or ""
+ a = row.get(asst_col) or ""
+ if u and a:
+ return [{"role": "user", "content": str(u)}, {"role": "assistant", "content": str(a)}]
+ return None
+
+ if row.get("messages"):
+ return _normalize_messages(row["messages"])
+
+ return None
+
+
+def _normalize_messages(msgs) -> list[dict] | None:
+ out = []
+ for m in msgs:
+ if not isinstance(m, dict):
+ continue
+ role = m.get("role") or "user"
+ content = m.get("content") or m.get("text") or ""
+ if content:
+ out.append({"role": role, "content": str(content)})
+ return out or None
diff --git a/scripts/requirements-train.txt b/scripts/requirements-train.txt
new file mode 100644
index 0000000..d48a70d
--- /dev/null
+++ b/scripts/requirements-train.txt
@@ -0,0 +1,9 @@
+torch>=2.1.0
+transformers>=4.40.0
+datasets>=2.18.0
+peft>=0.10.0
+bitsandbytes>=0.43.0
+trl>=0.8.0
+accelerate>=0.28.0
+sentencepiece>=0.2.0
+protobuf>=3.20.0
diff --git a/scripts/train_qlora.py b/scripts/train_qlora.py
index 9e8ee61..c795ac8 100644
--- a/scripts/train_qlora.py
+++ b/scripts/train_qlora.py
@@ -1,12 +1,16 @@
#!/usr/bin/env python3
-"""Minimal QLoRA trainer stub for Swarm Assistent.
-Requires: pip install torch transformers datasets peft bitsandbytes trl accelerate
-Configure runner in Assistent settings or replace with LLaMA-Factory CLI."""
+"""QLoRA SFT trainer for Swarm Assistent."""
+from __future__ import annotations
+
import argparse
import json
import os
import sys
-import time
+import traceback
+
+sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
+
+from hf_dataset_map import row_to_messages
def log(msg, log_path):
@@ -17,70 +21,226 @@ def log(msg, log_path):
f.write(line + "\n")
+class StepLogger:
+ def __init__(self, log_path, total_steps):
+ self.log_path = log_path
+ self.total_steps = max(1, total_steps)
+
+ def on_log(self, logs):
+ loss = logs.get("loss")
+ step = logs.get("step") or logs.get("global_step")
+ if step is None:
+ return
+ if loss is not None:
+ log(f"step {int(step)}/{self.total_steps} loss: {float(loss):.4f}", self.log_path)
+
+
+def detect_target_modules(model):
+ names = {n.split(".")[-1] for n, _ in model.named_modules()}
+ candidates = [
+ ["q_proj", "k_proj", "v_proj", "o_proj"],
+ ["q_proj", "v_proj"],
+ ["Wqkv", "out_proj"],
+ ["c_attn", "c_proj"],
+ ]
+ for group in candidates:
+ if all(g in names for g in group):
+ return group
+ return ["q_proj", "v_proj"]
+
+
+def load_sft_dataset(cfg, log_path):
+ from datasets import Dataset, load_dataset
+
+ hf_dataset = cfg.get("hf_dataset")
+ dataset_path = cfg.get("dataset_path")
+ max_samples = int(cfg.get("max_samples") or 0)
+ hf_mapping = cfg.get("hf_mapping") or {}
+ schema = cfg.get("hf_schema") or {}
+
+ rows = []
+
+ if hf_dataset:
+ log(f"loading HF dataset {hf_dataset}", log_path)
+ ds = load_dataset(hf_dataset, split="train")
+ if max_samples > 0:
+ ds = ds.select(range(min(len(ds), max_samples)))
+ for ex in ds:
+ msgs = row_to_messages(dict(ex), schema, hf_mapping)
+ if msgs:
+ rows.append({"messages": msgs})
+ elif dataset_path and os.path.isfile(dataset_path):
+ log(f"loading JSONL {dataset_path}", log_path)
+ with open(dataset_path, encoding="utf-8") as f:
+ for line in f:
+ line = line.strip()
+ if not line:
+ continue
+ try:
+ obj = json.loads(line)
+ except json.JSONDecodeError:
+ continue
+ msgs = obj.get("messages")
+ if msgs:
+ rows.append({"messages": msgs})
+ if max_samples > 0 and len(rows) >= max_samples:
+ break
+ else:
+ log("error: no dataset_path or hf_dataset", log_path)
+ sys.exit(1)
+
+ if not rows:
+ log("error: no training rows after mapping", log_path)
+ sys.exit(1)
+
+ log(f"dataset rows: {len(rows)}", log_path)
+ return Dataset.from_list(rows)
+
+
def main():
p = argparse.ArgumentParser()
p.add_argument("--config", required=True)
p.add_argument("--log", required=True)
args = p.parse_args()
+
with open(args.config, encoding="utf-8") as f:
cfg = json.load(f)
+
adapter_dir = cfg.get("adapter_dir", "adapter")
os.makedirs(adapter_dir, exist_ok=True)
- dataset_path = cfg.get("dataset_path")
- hf_dataset = cfg.get("hf_dataset")
- log(f"Swarm Assistent QLoRA stub starting base={cfg.get('base_model')}", args.log)
- if not dataset_path and not hf_dataset:
- log("error: no dataset", args.log)
- sys.exit(1)
+
+ log(f"Swarm Assistent QLoRA starting base={cfg.get('base_model')}", args.log)
+
try:
import torch
- from datasets import load_dataset
- from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer
- from peft import LoraConfig, get_peft_model, TaskType
+ from peft import LoraConfig, TaskType, get_peft_model
+ from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig, TrainerCallback, TrainingArguments
+ from trl import SFTTrainer
except ImportError as e:
- log(f"error: missing python deps ({e}). pip install torch transformers datasets peft bitsandbytes trl accelerate", args.log)
+ log(f"error: missing python deps ({e}). pip install -r scripts/requirements-train.txt", args.log)
sys.exit(2)
+
base = cfg.get("base_model")
+ if not base:
+ log("error: base_model required", args.log)
+ sys.exit(1)
+
+ try:
+ ds = load_sft_dataset(cfg, args.log)
+ except Exception as e:
+ log(f"error: dataset load failed: {e}", args.log)
+ traceback.print_exc()
+ sys.exit(1)
+
+ four_bit = bool(cfg.get("four_bit", True))
+ seq_len = int(cfg.get("seq_len") or 2048)
+ rank = int(cfg.get("rank") or 16)
+ alpha = int(cfg.get("alpha") or 32)
+ lr = float(cfg.get("lr") or 2e-4)
+ epochs = int(cfg.get("epochs") or 3)
+ batch_size = int(cfg.get("batch_size") or 1)
+ grad_accum = int(cfg.get("gradient_accumulation_steps") or 4)
+
log(f"loading model {base}", args.log)
tokenizer = AutoTokenizer.from_pretrained(base, trust_remote_code=True)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
+
+ bnb_config = None
+ if four_bit:
+ bnb_config = BitsAndBytesConfig(
+ load_in_4bit=True,
+ bnb_4bit_quant_type="nf4",
+ bnb_4bit_compute_dtype=torch.float16,
+ bnb_4bit_use_double_quant=True,
+ )
+
model = AutoModelForCausalLM.from_pretrained(
base,
- load_in_4bit=bool(cfg.get("four_bit", True)),
+ quantization_config=bnb_config,
device_map="auto",
trust_remote_code=True,
+ torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,
)
+
+ target_modules = detect_target_modules(model)
+ log(f"lora target_modules: {target_modules}", args.log)
+
lora = LoraConfig(
- r=int(cfg.get("rank", 16)),
- lora_alpha=int(cfg.get("alpha", 32)),
+ r=rank,
+ lora_alpha=alpha,
+ lora_dropout=0.05,
task_type=TaskType.CAUSAL_LM,
- target_modules=["q_proj", "v_proj", "k_proj", "o_proj"],
+ target_modules=target_modules,
)
model = get_peft_model(model, lora)
- if hf_dataset:
- ds = load_dataset(hf_dataset, split="train")
- else:
- ds = load_dataset("json", data_files=dataset_path, split="train")
- def fmt(ex):
- msgs = ex.get("messages")
- if msgs:
- text = tokenizer.apply_chat_template(msgs, tokenize=False)
- else:
- text = ex.get("text") or ""
- return {"text": text}
+ def formatting_func(examples):
+ texts = []
+ for msgs in examples["messages"]:
+ try:
+ text = tokenizer.apply_chat_template(msgs, tokenize=False, add_generation_prompt=False)
+ except Exception:
+ parts = []
+ for m in msgs:
+ role = m.get("role", "user")
+ content = m.get("content", "")
+ parts.append(f"{role}: {content}")
+ text = "\n".join(parts)
+ texts.append(text)
+ return texts
+
+ total_steps = max(1, (len(ds) * epochs) // max(1, batch_size * grad_accum))
+ log(f"planned steps ~{total_steps}", args.log)
+
+ training_args = TrainingArguments(
+ output_dir=adapter_dir,
+ num_train_epochs=epochs,
+ per_device_train_batch_size=batch_size,
+ gradient_accumulation_steps=grad_accum,
+ learning_rate=lr,
+ logging_steps=1,
+ save_steps=max(50, total_steps // 10),
+ save_total_limit=2,
+ fp16=torch.cuda.is_available(),
+ bf16=False,
+ report_to="none",
+ remove_unused_columns=False,
+ max_grad_norm=0.3,
+ warmup_ratio=0.03,
+ lr_scheduler_type="cosine",
+ )
+
+ step_logger = StepLogger(args.log, total_steps)
+
+ class LossCallback(TrainerCallback):
+ def on_log(self, args_, state, control, logs=None, **kwargs):
+ if logs:
+ step_logger.on_log({**logs, "step": state.global_step})
+
+ try:
+ trainer = SFTTrainer(
+ model=model,
+ args=training_args,
+ train_dataset=ds,
+ tokenizer=tokenizer,
+ formatting_func=formatting_func,
+ max_seq_length=seq_len,
+ packing=False,
+ callbacks=[LossCallback()],
+ )
+ log("training started", args.log)
+ trainer.train()
+ trainer.save_model(adapter_dir)
+ tokenizer.save_pretrained(adapter_dir)
+ except Exception as e:
+ log(f"error: training failed: {e}", args.log)
+ traceback.print_exc()
+ sys.exit(3)
- ds = ds.map(fmt)
- epochs = int(cfg.get("epochs", 3))
- steps = max(1, min(len(ds), 100) * epochs)
- for i in range(1, steps + 1):
- log(f"step {i}/{steps} loss: {1.0 / i:.4f}", args.log)
- time.sleep(0.05)
- model.save_pretrained(adapter_dir)
- tokenizer.save_pretrained(adapter_dir)
with open(os.path.join(adapter_dir, "train_done.json"), "w", encoding="utf-8") as f:
- json.dump({"ok": True, "base": base}, f)
+ json.dump({"ok": True, "base": base, "rows": len(ds)}, f)
+
log("training complete", args.log)
diff --git a/src/training.js b/src/training.js
index b94d15c..cb39e0a 100644
--- a/src/training.js
+++ b/src/training.js
@@ -17,6 +17,7 @@ export function attachTraining(SA) {
hfResults: [],
hfSelected: null,
hfCheck: null,
+ hfMapping: null,
trainWs: null,
polling: null,
agentSettings: { enabled: true, auto_link_on_approve: true, heard_quota: 3 },
@@ -97,7 +98,10 @@ export function attachTraining(SA) {
refreshSamples();
loadAgentHeardSettings();
}
- if (state.ttab === 'train') syncModelfileModels();
+ if (state.ttab === 'train') {
+ syncModelfileModels();
+ syncQloraModels();
+ }
if (state.ttab === 'models') refreshTrainModels();
}
@@ -162,6 +166,81 @@ export function attachTraining(SA) {
await refreshSamples();
}
+ function hfStringColumns(check) {
+ const cols = check?.schema?.columns;
+ if (Array.isArray(cols) && cols.length) return cols;
+ const feats = check?.features;
+ if (Array.isArray(feats)) {
+ return feats.map((f) => f?.name).filter(Boolean);
+ }
+ if (feats && typeof feats === 'object') return Object.keys(feats);
+ return [];
+ }
+
+ function renderHfMappingUI(check) {
+ const row = $('sa_hf_mapping_row');
+ if (!row) return;
+ const gate = check?.gate;
+ const schemaKind = check?.schema?.kind;
+ const needsMapping = gate === 'mapping' || schemaKind === 'fiction_tags_text';
+ row.hidden = !needsMapping;
+ if (!needsMapping) {
+ state.hfMapping = null;
+ return;
+ }
+ const cols = hfStringColumns(check);
+ const userSel = $('sa_hf_user_col');
+ const asstSel = $('sa_hf_asst_col');
+ const presetSel = $('sa_hf_mapping_preset');
+ if (userSel) {
+ userSel.innerHTML = cols.map((c) => ``).join('');
+ if (cols.includes('tags')) userSel.value = 'tags';
+ else if (cols.includes('title')) userSel.value = 'title';
+ }
+ if (asstSel) {
+ asstSel.innerHTML = cols.map((c) => ``).join('');
+ if (cols.includes('text')) asstSel.value = 'text';
+ else if (cols.includes('output')) asstSel.value = 'output';
+ }
+ if (schemaKind === 'fiction_tags_text' && presetSel) {
+ presetSel.value = 'fiction_tags_text';
+ state.hfMapping = { kind: 'fiction_tags_text', preset: 'fiction_tags_text' };
+ }
+ }
+
+ function buildHfMappingPayload() {
+ const preset = $('sa_hf_mapping_preset')?.value;
+ if (preset === 'fiction_tags_text') {
+ return { kind: 'fiction_tags_text', preset: 'fiction_tags_text' };
+ }
+ const userCol = $('sa_hf_user_col')?.value;
+ const asstCol = $('sa_hf_asst_col')?.value;
+ if (userCol && asstCol) {
+ return { kind: 'custom', user_col: userCol, assistant_col: asstCol };
+ }
+ return state.hfMapping;
+ }
+
+ async function syncQloraModels() {
+ try {
+ const baseUrl = $('sa_base_url')?.value || localStorage.getItem('swarm_assistent_base_url') || '';
+ const data = await SA.request('AssistentListModels', { baseUrl });
+ const models = data?.models || [];
+ const sel = $('sa_qlora_ollama_base');
+ if (!sel) return;
+ const cur = sel.value;
+ sel.innerHTML = '';
+ for (const m of models) {
+ const opt = document.createElement('option');
+ opt.value = m;
+ opt.textContent = m;
+ sel.appendChild(opt);
+ }
+ if (cur) sel.value = cur;
+ else if ($('sa_model')?.value) sel.value = $('sa_model').value;
+ } catch (e) { /* ignore */ }
+ }
+
function renderHfList() {
const root = $('sa_hf_list');
if (!root) return;
@@ -216,6 +295,7 @@ export function attachTraining(SA) {
}
const importRow = $('sa_hf_import_row');
if (importRow) importRow.hidden = data.gate === 'rejected';
+ renderHfMappingUI(data);
} catch (e) {
if (status) status.textContent = String(e.message || e);
}
@@ -228,8 +308,9 @@ export function attachTraining(SA) {
}
const id = state.hfSelected || state.hfCheck.id;
const limit = Number($('sa_hf_import_limit')?.value) || 200;
+ const mapping = buildHfMappingPayload();
try {
- const data = await SA.request('AssistentImportHfDataset', { dataset: id, limit });
+ const data = await SA.request('AssistentImportHfDataset', { dataset: id, limit, mapping });
setTrainStatus(`Импортировано: ${data.imported}${data.runner_only ? ' (runner-only)' : ''}`);
await refreshSamples();
} catch (e) {
@@ -303,8 +384,9 @@ export function attachTraining(SA) {
async function pollTrainJob() {
try {
const data = await SA.request('AssistentGetTrainJob', {});
- const prog = data?.job?.progress_json ? JSON.parse(data.job.progress_json) : null;
+ const prog = data?.progress || (data?.job?.progress_json ? JSON.parse(data.job.progress_json) : null);
const active = data?.training_active || data?.job?.status === 'running';
+ const status = data?.job?.status || prog?.status;
setTrainingLock(active, prog?.status === 'running' ? `Тренировка · ${prog?.percent ?? 0}%` : 'Идёт тренировка…');
const logEl = $('sa_train_log');
const bar = $('sa_train_progress_fill');
@@ -318,6 +400,24 @@ export function attachTraining(SA) {
clearInterval(state.polling);
state.polling = null;
$('sa_btn_qlora_cancel').hidden = true;
+ setTrainingLock(false);
+ if (status === 'completed' || status === 'completed_with_warnings') {
+ const ollama = prog?.ollama;
+ if (ollama?.success) {
+ setTrainStatus(`Готово: модель ${ollama.name} в Ollama`);
+ SA.app?.refreshModels?.();
+ } else if (ollama?.skipped) {
+ setTrainStatus(ollama.note || ollama.error || 'Адаптер сохранён, Ollama — вручную');
+ } else if (ollama?.error) {
+ setTrainStatus(`Обучение OK, Ollama: ${ollama.error}`);
+ } else if (status === 'completed_with_warnings') {
+ setTrainStatus('Обучение завершено с предупреждениями — см. лог');
+ } else {
+ setTrainStatus('QLoRA завершено');
+ }
+ } else if (status === 'failed') {
+ setTrainStatus(`Ошибка тренировки (exit ${prog?.exit_code ?? '?'})`);
+ }
}
} catch (e) { /* ignore */ }
}
@@ -325,18 +425,23 @@ export function attachTraining(SA) {
async function startQlora() {
setTrainStatus('Запуск…');
try {
+ const hfDs = ($('sa_qlora_hf_dataset')?.value || '').trim();
+ const mapping = hfDs ? buildHfMappingPayload() : undefined;
await SA.request('AssistentStartTrainJob', {
base_url: $('sa_base_url')?.value,
chat_model: $('sa_model')?.value,
base_model: $('sa_qlora_base')?.value,
+ ollama_base: $('sa_qlora_ollama_base')?.value,
output_name: $('sa_qlora_name')?.value,
rank: Number($('sa_qlora_rank')?.value) || 16,
alpha: Number($('sa_qlora_alpha')?.value) || 32,
lr: Number($('sa_qlora_lr')?.value) || 0.0002,
epochs: Number($('sa_qlora_epochs')?.value) || 3,
seq_len: Number($('sa_qlora_seq')?.value) || 2048,
+ max_samples: Number($('sa_qlora_max_samples')?.value) || 0,
four_bit: !!$('sa_qlora_4bit')?.checked,
- hf_dataset: ($('sa_qlora_hf_dataset')?.value || '').trim() || undefined,
+ hf_dataset: hfDs || undefined,
+ hf_mapping: mapping,
});
$('sa_btn_qlora_cancel').hidden = false;
setTrainingLock(true, 'Идёт тренировка…');
@@ -377,10 +482,12 @@ export function attachTraining(SA) {
try {
await SA.request('AssistentSaveRunnerSettings', {
python: $('sa_runner_python')?.value,
- kind: $('sa_runner_kind')?.value,
+ kind: $('sa_runner_kind')?.value || 'builtin',
workdir: $('sa_runner_workdir')?.value,
cmd: $('sa_runner_cmd')?.value,
gguf_script: $('sa_runner_gguf_script')?.value,
+ gguf_base_path: $('sa_runner_gguf_base')?.value,
+ gguf_cmd: $('sa_runner_gguf_cmd')?.value,
});
setTrainStatus('Раннер сохранён');
} catch (e) {
@@ -393,10 +500,12 @@ export function attachTraining(SA) {
const data = await SA.request('AssistentGetRunnerSettings', {});
const s = data?.settings || {};
if ($('sa_runner_python') && s.python) $('sa_runner_python').value = s.python;
- if ($('sa_runner_kind') && s.kind) $('sa_runner_kind').value = s.kind;
+ if ($('sa_runner_kind')) $('sa_runner_kind').value = s.kind || 'builtin';
if ($('sa_runner_workdir') && s.workdir) $('sa_runner_workdir').value = s.workdir;
if ($('sa_runner_cmd') && s.cmd) $('sa_runner_cmd').value = s.cmd;
if ($('sa_runner_gguf_script') && s.gguf_script) $('sa_runner_gguf_script').value = s.gguf_script;
+ if ($('sa_runner_gguf_base') && s.gguf_base_path) $('sa_runner_gguf_base').value = s.gguf_base_path;
+ if ($('sa_runner_gguf_cmd') && s.gguf_cmd) $('sa_runner_gguf_cmd').value = s.gguf_cmd;
} catch (e) { /* ignore */ }
}
@@ -491,6 +600,9 @@ export function attachTraining(SA) {
await checkHfLink();
});
$('sa_btn_hf_check')?.addEventListener('click', checkHfLink);
+ $('sa_hf_mapping_preset')?.addEventListener('change', () => {
+ state.hfMapping = buildHfMappingPayload();
+ });
$('sa_btn_hf_import')?.addEventListener('click', importHf);
document.querySelectorAll('input[name="sa_train_mode"]').forEach((r) => {
r.addEventListener('change', () => setTrainMode(r.value));