From 474a674a3562b7cda6984e75f8d380ba5a4c4dd1 Mon Sep 17 00:00:00 2001 From: Leonid Pershin Date: Sat, 22 Aug 2026 14:45:07 +0300 Subject: [PATCH] Ship Assistent 0.13.0: real QLoRA pipeline and GGUF Ollama register. Replace fake train loop with TRL SFTTrainer, HF column mapping with fiction preset, safetensors to GGUF conversion, and ollama create using ollama_base plus ADAPTER. Co-authored-by: Cursor --- Assets/assistent.bundle.js | 123 ++++++++++++- Assets/assistent.css | 15 ++ AssistentHuggingFace.cs | 153 +++++++++++++++- AssistentTraining.cs | 11 +- AssistentTrainingJobs.cs | 322 ++++++++++++++++++++++++++------- README.md | 12 ++ SwarmAssistentExtension.cs | 6 +- Tabs/Text2Image/Assistent.html | 30 ++- scripts/hf_dataset_map.py | 94 ++++++++++ scripts/requirements-train.txt | 9 + scripts/train_qlora.py | 236 ++++++++++++++++++++---- src/training.js | 124 ++++++++++++- 12 files changed, 1002 insertions(+), 133 deletions(-) create mode 100644 scripts/hf_dataset_map.py create mode 100644 scripts/requirements-train.txt diff --git a/Assets/assistent.bundle.js b/Assets/assistent.bundle.js index 0b966de..1411677 100644 --- a/Assets/assistent.bundle.js +++ b/Assets/assistent.bundle.js @@ -9168,6 +9168,7 @@ ${HELP_TEXT}`); hfResults: [], hfSelected: null, hfCheck: null, + hfMapping: null, trainWs: null, polling: null, agentSettings: { enabled: true, auto_link_on_approve: true, heard_quota: 3 }, @@ -9242,7 +9243,10 @@ ${HELP_TEXT}`); refreshSamples(); loadAgentHeardSettings(); } - if (state.ttab === "train") syncModelfileModels(); + if (state.ttab === "train") { + syncModelfileModels(); + syncQloraModels(); + } if (state.ttab === "models") refreshTrainModels(); } async function refreshSamples() { @@ -9303,6 +9307,78 @@ ${HELP_TEXT}`); await SA2.request("AssistentUpsertTrainSample", patch); 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 SA2.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) { + } + } function renderHfList() { const root = $("sa_hf_list"); if (!root) return; @@ -9353,6 +9429,7 @@ ${HELP_TEXT}`); } 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); } @@ -9364,8 +9441,9 @@ ${HELP_TEXT}`); } const id = state.hfSelected || state.hfCheck.id; const limit = Number($("sa_hf_import_limit")?.value) || 200; + const mapping = buildHfMappingPayload(); try { - const data = await SA2.request("AssistentImportHfDataset", { dataset: id, limit }); + const data = await SA2.request("AssistentImportHfDataset", { dataset: id, limit, mapping }); setTrainStatus(`\u0418\u043C\u043F\u043E\u0440\u0442\u0438\u0440\u043E\u0432\u0430\u043D\u043E: ${data.imported}${data.runner_only ? " (runner-only)" : ""}`); await refreshSamples(); } catch (e) { @@ -9435,8 +9513,9 @@ ${HELP_TEXT}`); async function pollTrainJob() { try { const data = await SA2.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" ? `\u0422\u0440\u0435\u043D\u0438\u0440\u043E\u0432\u043A\u0430 \xB7 ${prog?.percent ?? 0}%` : "\u0418\u0434\u0451\u0442 \u0442\u0440\u0435\u043D\u0438\u0440\u043E\u0432\u043A\u0430\u2026"); const logEl = $("sa_train_log"); const bar = $("sa_train_progress_fill"); @@ -9450,6 +9529,24 @@ ${HELP_TEXT}`); 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(`\u0413\u043E\u0442\u043E\u0432\u043E: \u043C\u043E\u0434\u0435\u043B\u044C ${ollama.name} \u0432 Ollama`); + SA2.app?.refreshModels?.(); + } else if (ollama?.skipped) { + setTrainStatus(ollama.note || ollama.error || "\u0410\u0434\u0430\u043F\u0442\u0435\u0440 \u0441\u043E\u0445\u0440\u0430\u043D\u0451\u043D, Ollama \u2014 \u0432\u0440\u0443\u0447\u043D\u0443\u044E"); + } else if (ollama?.error) { + setTrainStatus(`\u041E\u0431\u0443\u0447\u0435\u043D\u0438\u0435 OK, Ollama: ${ollama.error}`); + } else if (status === "completed_with_warnings") { + setTrainStatus("\u041E\u0431\u0443\u0447\u0435\u043D\u0438\u0435 \u0437\u0430\u0432\u0435\u0440\u0448\u0435\u043D\u043E \u0441 \u043F\u0440\u0435\u0434\u0443\u043F\u0440\u0435\u0436\u0434\u0435\u043D\u0438\u044F\u043C\u0438 \u2014 \u0441\u043C. \u043B\u043E\u0433"); + } else { + setTrainStatus("QLoRA \u0437\u0430\u0432\u0435\u0440\u0448\u0435\u043D\u043E"); + } + } else if (status === "failed") { + setTrainStatus(`\u041E\u0448\u0438\u0431\u043A\u0430 \u0442\u0440\u0435\u043D\u0438\u0440\u043E\u0432\u043A\u0438 (exit ${prog?.exit_code ?? "?"})`); + } } } catch (e) { } @@ -9457,18 +9554,23 @@ ${HELP_TEXT}`); async function startQlora() { setTrainStatus("\u0417\u0430\u043F\u0443\u0441\u043A\u2026"); try { + const hfDs = ($("sa_qlora_hf_dataset")?.value || "").trim(); + const mapping = hfDs ? buildHfMappingPayload() : void 0; await SA2.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) || 2e-4, 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() || void 0 + hf_dataset: hfDs || void 0, + hf_mapping: mapping }); $("sa_btn_qlora_cancel").hidden = false; setTrainingLock(true, "\u0418\u0434\u0451\u0442 \u0442\u0440\u0435\u043D\u0438\u0440\u043E\u0432\u043A\u0430\u2026"); @@ -9504,10 +9606,12 @@ ${HELP_TEXT}`); try { await SA2.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_script: $("sa_runner_gguf_script")?.value, + gguf_base_path: $("sa_runner_gguf_base")?.value, + gguf_cmd: $("sa_runner_gguf_cmd")?.value }); setTrainStatus("\u0420\u0430\u043D\u043D\u0435\u0440 \u0441\u043E\u0445\u0440\u0430\u043D\u0451\u043D"); } catch (e) { @@ -9519,10 +9623,12 @@ ${HELP_TEXT}`); const data = await SA2.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) { } } @@ -9629,6 +9735,9 @@ ${HELP_TEXT}`); 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)); diff --git a/Assets/assistent.css b/Assets/assistent.css index 32f691e..858c162 100644 --- a/Assets/assistent.css +++ b/Assets/assistent.css @@ -2281,6 +2281,21 @@ width: 2.75rem; } +.sa-hf-mapping-row { + display: flex; + flex-wrap: wrap; + gap: 0.45rem 0.75rem; + align-items: center; + padding: 0.35rem 0; +} + +.sa-hf-mapping-row label { + display: inline-flex; + align-items: center; + gap: 0.3rem; + font-size: 0.82rem; +} + .sa-hf-panel { border: 1px solid color-mix(in srgb, currentColor 18%, transparent); border-radius: 0.45rem; diff --git a/AssistentHuggingFace.cs b/AssistentHuggingFace.cs index e75dd5d..58ce218 100644 --- a/AssistentHuggingFace.cs +++ b/AssistentHuggingFace.cs @@ -188,8 +188,14 @@ public partial class SwarmAssistentExtension return CacheHfCheck(cacheKey, result); } JObject rowsData = JObject.Parse(rowsBody); - JObject features = rowsData["features"] as JObject; + JObject features = FeaturesToObject(rowsData["features"]); (string gate, string reason, JObject schema) = ClassifyHfFeatures(features); + if (gate == "mapping" && TryFictionTagsTextPreset(features, out JObject presetSchema)) + { + gate = "ok"; + reason = "Fiction preset: title/tags → user, text → assistant"; + schema = presetSchema; + } result["gate"] = gate; result["reason"] = reason; result["schema"] = schema; @@ -197,7 +203,7 @@ public partial class SwarmAssistentExtension result["split"] = split; result["features"] = features; result["sample_rows"] = rowsData["rows"]; - result["runner_only"] = HasHugeSizeTag(datasetId); + result["runner_only"] = await HasHugeSizeTagAsync(session, datasetId); return CacheHfCheck(cacheKey, result); } catch (Exception ex) @@ -207,7 +213,89 @@ public partial class SwarmAssistentExtension } } - static bool HasHugeSizeTag(string datasetId) => false; + static bool TryFictionTagsTextPreset(JObject features, out JObject schema) + { + schema = null; + if (features is null) + { + return false; + } + HashSet names = new(StringComparer.OrdinalIgnoreCase); + foreach (JProperty p in features.Properties()) + { + names.Add(p.Name); + } + if (!names.Contains("text") || (!names.Contains("tags") && !names.Contains("title"))) + { + return false; + } + schema = new JObject + { + ["kind"] = "fiction_tags_text", + ["assistant_col"] = "text", + }; + return true; + } + + async Task HasHugeSizeTagAsync(Session session, string datasetId) + { + try + { + using HttpRequestMessage req = HfRequest($"{HfHubApi}/{Uri.EscapeDataString(datasetId)}", session); + using HttpResponseMessage resp = await HttpClient.SendAsync(req); + if (!resp.IsSuccessStatusCode) + { + return false; + } + JObject meta = JObject.Parse(await resp.Content.ReadAsStringAsync()); + foreach (JToken t in meta["tags"] as JArray ?? []) + { + string tag = t?.ToString() ?? ""; + if (tag.Contains("100K<", StringComparison.OrdinalIgnoreCase) + || tag.Contains("1M<", StringComparison.OrdinalIgnoreCase) + || tag.Contains("10M<", StringComparison.OrdinalIgnoreCase) + || tag.Contains("100M<", StringComparison.OrdinalIgnoreCase)) + { + return true; + } + } + } + catch (Exception ex) + { + Logs.Debug($"HasHugeSizeTag {datasetId}: {ex.Message}"); + } + return false; + } + + static JObject ResolveHfMapping(JObject check, JObject mapping) + { + JObject schema = check?["schema"] as JObject; + string kind = schema?["kind"]?.ToString(); + if (string.Equals(kind, "fiction_tags_text", StringComparison.OrdinalIgnoreCase)) + { + return new JObject { ["kind"] = "fiction_tags_text", ["preset"] = "fiction_tags_text" }; + } + if (mapping is not null && mapping.Count > 0) + { + return mapping; + } + if (string.Equals(mapping?["preset"]?.ToString(), "fiction_tags_text", StringComparison.OrdinalIgnoreCase)) + { + return new JObject { ["kind"] = "fiction_tags_text" }; + } + return mapping; + } + + static bool MappingRequired(JObject check, JObject mapping) + { + string gate = check?["gate"]?.ToString(); + if (gate != "mapping") + { + return false; + } + JObject resolved = ResolveHfMapping(check, mapping); + return resolved is null || !resolved.Properties().Any(); + } JObject CacheHfCheck(string cacheKey, JObject result) { @@ -225,6 +313,33 @@ public partial class SwarmAssistentExtension return result; } + static JObject FeaturesToObject(JToken tok) + { + if (tok is JObject obj) + { + return obj; + } + if (tok is JArray arr) + { + JObject map = new(); + foreach (JToken t in arr) + { + if (t is not JObject row) + { + continue; + } + string name = row["name"]?.ToString(); + if (string.IsNullOrWhiteSpace(name)) + { + continue; + } + map[name] = row["type"] ?? row; + } + return map.Count > 0 ? map : null; + } + return null; + } + static (string gate, string reason, JObject schema) ClassifyHfFeatures(JObject features) { if (features is null || !features.Properties().Any()) @@ -325,10 +440,11 @@ public partial class SwarmAssistentExtension { return new JObject { ["error"] = check["reason"]?.ToString() ?? "rejected" }; } - if (gate == "mapping" && (mapping is null || mapping.Count == 0)) + if (MappingRequired(check, mapping)) { return new JObject { ["error"] = "Нужен маппинг колонок", ["check"] = check }; } + mapping = ResolveHfMapping(check, mapping); int take = Math.Clamp(limit, 1, 5000); JArray rows = []; string config = check["config"]?.ToString() ?? "default"; @@ -445,6 +561,35 @@ public partial class SwarmAssistentExtension new JObject { ["role"] = "assistant", ["content"] = row["answer"]?.ToString() ?? "" }, }; } + if (kind == "fiction_tags_text") + { + string text = row["text"]?.ToString() ?? ""; + if (string.IsNullOrWhiteSpace(text)) + { + return null; + } + List userParts = []; + string title = row["title"]?.ToString()?.Trim(); + string tags = row["tags"]?.ToString()?.Trim(); + if (!string.IsNullOrWhiteSpace(title)) + { + userParts.Add($"Title: {title}"); + } + if (!string.IsNullOrWhiteSpace(tags)) + { + userParts.Add($"Tags: {tags}"); + } + string user = userParts.Count > 0 ? string.Join("\n", userParts) : tags ?? title ?? ""; + if (string.IsNullOrWhiteSpace(user)) + { + return null; + } + return new JArray + { + new JObject { ["role"] = "user", ["content"] = user }, + new JObject { ["role"] = "assistant", ["content"] = text }, + }; + } if (kind == "custom" && mapping is not null) { string userCol = mapping["user_col"]?.ToString(); diff --git a/AssistentTraining.cs b/AssistentTraining.cs index 38feac1..35db077 100644 --- a/AssistentTraining.cs +++ b/AssistentTraining.cs @@ -296,6 +296,7 @@ public partial class SwarmAssistentExtension { List samples = Memory.ListTrainSamples(status, null, null, 5000); StringBuilder sb = new(); + int exported = 0; foreach (JObject s in samples) { JArray messages = s["messages"] as JArray ?? []; @@ -303,6 +304,7 @@ public partial class SwarmAssistentExtension { continue; } + exported++; if (string.Equals(format, "sharegpt", StringComparison.OrdinalIgnoreCase)) { JArray conv = []; @@ -328,7 +330,7 @@ public partial class SwarmAssistentExtension } string path = Path.Combine(TrainingRoot(), "datasets", $"export_{DateTimeOffset.UtcNow.ToUnixTimeMilliseconds()}.jsonl"); await File.WriteAllTextAsync(path, sb.ToString(), Encoding.UTF8); - return new JObject { ["success"] = true, ["path"] = path, ["count"] = samples.Count, ["content"] = sb.ToString() }; + return new JObject { ["success"] = true, ["path"] = path, ["count"] = exported, ["content"] = sb.ToString() }; } catch (Exception ex) { @@ -416,11 +418,18 @@ public partial class SwarmAssistentExtension { await Task.CompletedTask; JObject job = string.IsNullOrWhiteSpace(id) ? Memory.GetActiveTrainJob() : Memory.GetTrainJob(id); + if (job is not null && TrainingJobManager.IsRunning && string.Equals(job["id"]?.ToString(), TrainingJobManager.CurrentJobId, StringComparison.OrdinalIgnoreCase)) + { + JObject live = TrainingJobManager.GetProgress(); + job["progress_json"] = live.ToString(Newtonsoft.Json.Formatting.None); + job["status"] = live["status"]?.ToString() ?? job["status"]; + } return new JObject { ["success"] = true, ["job"] = job, ["training_active"] = TrainingJobManager.IsRunning, + ["progress"] = TrainingJobManager.IsRunning ? TrainingJobManager.GetProgress() : null, }; } diff --git a/AssistentTrainingJobs.cs b/AssistentTrainingJobs.cs index 0b0ae18..3fce3ea 100644 --- a/AssistentTrainingJobs.cs +++ b/AssistentTrainingJobs.cs @@ -35,30 +35,32 @@ public partial class SwarmAssistentExtension { return new JObject { ["error"] = "base_model and output_name required" }; } + string hfDataset = null; + JObject hfCheck = null; + JObject hfMapping = raw["hf_mapping"] as JObject; if (!string.IsNullOrWhiteSpace(raw["hf_dataset"]?.ToString())) { - string dsId = NormalizeHfDatasetId(raw["hf_dataset"]?.ToString()); - if (dsId is null) + hfDataset = NormalizeHfDatasetId(raw["hf_dataset"]?.ToString()); + if (hfDataset is null) { return new JObject { ["error"] = "invalid hf_dataset id" }; } - JObject check = await CheckHfDatasetInternal(session, dsId, useCache: true); - if (check["gate"]?.ToString() == "rejected") + hfCheck = await CheckHfDatasetInternal(session, hfDataset, useCache: true); + if (hfCheck["gate"]?.ToString() == "rejected") { - return new JObject { ["error"] = check["reason"]?.ToString() ?? "hf dataset rejected" }; + return new JObject { ["error"] = hfCheck["reason"]?.ToString() ?? "hf dataset rejected" }; + } + hfMapping = ResolveHfMapping(hfCheck, hfMapping); + if (MappingRequired(hfCheck, hfMapping)) + { + return new JObject { ["error"] = "Нужен маппинг колонок для HF набора", ["check"] = hfCheck }; } - raw["hf_dataset"] = dsId; } JObject runner = Config.LoadTrainingRunner(); - string python = runner["python"]?.ToString()?.Trim(); - if (string.IsNullOrWhiteSpace(python)) - { - python = "python"; - } string kind = runner["kind"]?.ToString()?.Trim(); if (string.IsNullOrWhiteSpace(kind)) { - return new JObject { ["error"] = "QLoRA-раннер не настроен (Настройки → Модели)" }; + kind = "builtin"; } string baseUrl = NormalizeBaseUrl(raw["base_url"]?.ToString()); string chatModel = raw["chat_model"]?.ToString()?.Trim(); @@ -66,11 +68,26 @@ public partial class SwarmAssistentExtension { await AssistentParkLlm(session, baseUrl, chatModel); } - JObject export = await AssistentExportDataset(session, "approved", "jsonl"); - string datasetPath = export["path"]?.ToString(); - if (string.IsNullOrWhiteSpace(datasetPath) || !File.Exists(datasetPath)) + string datasetPath = null; + int exportCount = 0; + if (string.IsNullOrWhiteSpace(hfDataset)) { - return new JObject { ["error"] = "Нет одобренных примеров для тренировки" }; + JObject export = await AssistentExportDataset(session, "approved", "jsonl"); + datasetPath = export["path"]?.ToString(); + exportCount = export["count"]?.Value() ?? 0; + if (string.IsNullOrWhiteSpace(datasetPath) || !File.Exists(datasetPath) || exportCount <= 0) + { + return new JObject { ["error"] = "Нет одобренных примеров для тренировки (или укажи hf_dataset)" }; + } + } + else + { + JObject export = await AssistentExportDataset(session, "approved", "jsonl"); + exportCount = export["count"]?.Value() ?? 0; + if (exportCount > 0 && File.Exists(export["path"]?.ToString() ?? "")) + { + datasetPath = export["path"]?.ToString(); + } } long now = DateTimeOffset.UtcNow.ToUnixTimeMilliseconds(); string jobId = $"tj_{now}"; @@ -78,19 +95,29 @@ public partial class SwarmAssistentExtension Directory.CreateDirectory(jobDir); string configPath = Path.Combine(jobDir, "config.json"); string logPath = Path.Combine(jobDir, "log.txt"); + string adapterDir = Path.Combine(TrainingRoot(), "adapters", SanitizeAdapterName(outputName)); + Directory.CreateDirectory(adapterDir); JObject jobConfig = new() { ["base_model"] = hfBase, ["output_name"] = outputName, + ["ollama_base"] = raw["ollama_base"]?.ToString()?.Trim(), + ["gguf_base_path"] = raw["gguf_base_path"]?.ToString()?.Trim() ?? runner["gguf_base_path"]?.ToString()?.Trim(), ["dataset_path"] = datasetPath, - ["hf_dataset"] = raw["hf_dataset"], + ["hf_dataset"] = hfDataset, + ["hf_mapping"] = hfMapping, + ["hf_schema"] = hfCheck?["schema"], + ["max_samples"] = raw["max_samples"] ?? 0, ["rank"] = raw["rank"] ?? 16, ["alpha"] = raw["alpha"] ?? 32, ["lr"] = raw["lr"] ?? 0.0002, ["epochs"] = raw["epochs"] ?? 3, ["seq_len"] = raw["seq_len"] ?? 2048, ["four_bit"] = raw["four_bit"] ?? true, - ["adapter_dir"] = Path.Combine(TrainingRoot(), "adapters", outputName), + ["batch_size"] = raw["batch_size"] ?? 1, + ["gradient_accumulation_steps"] = raw["gradient_accumulation_steps"] ?? 4, + ["adapter_dir"] = adapterDir, + ["local_export_count"] = exportCount, }; await File.WriteAllTextAsync(configPath, jobConfig.ToString(Newtonsoft.Json.Formatting.Indented), Encoding.UTF8); Memory.SaveTrainJob(new JObject @@ -114,6 +141,21 @@ public partial class SwarmAssistentExtension return new JObject { ["success"] = true, ["job_id"] = jobId, ["log_path"] = logPath }; } + static string SanitizeAdapterName(string name) + { + if (string.IsNullOrWhiteSpace(name)) + { + return "adapter"; + } + char[] bad = Path.GetInvalidFileNameChars(); + StringBuilder sb = new(); + foreach (char c in name) + { + sb.Append(Array.IndexOf(bad, c) >= 0 ? '_' : c); + } + return sb.ToString().Trim(); + } + static string BuildRunnerCommand(JObject runner, string configPath, string logPath, string workDir) { string python = runner["python"]?.ToString()?.Trim(); @@ -121,10 +163,10 @@ public partial class SwarmAssistentExtension { python = "python"; } - string kind = runner["kind"]?.ToString()?.Trim() ?? "custom"; + string kind = runner["kind"]?.ToString()?.Trim() ?? "builtin"; string custom = runner["cmd"]?.ToString()?.Trim(); string scriptPath = Path.Combine(FilePath, "scripts", "train_qlora.py"); - if (kind == "custom" && !string.IsNullOrWhiteSpace(custom)) + if (string.Equals(kind, "custom", StringComparison.OrdinalIgnoreCase) && !string.IsNullOrWhiteSpace(custom)) { return custom .Replace("{python}", python, StringComparison.OrdinalIgnoreCase) @@ -175,27 +217,38 @@ public partial class SwarmAssistentExtension return new JObject { ["success"] = true }; } - internal async Task FinishTrainJobAsync(string jobId, bool success, string logPath, Session session, string baseUrl, string chatModel, string adapterDir, string outputName, string ggufScript) + internal async Task FinishTrainJobAsync( + string jobId, + bool success, + string logPath, + Session session, + string baseUrl, + string chatModel, + JObject jobConfig, + JObject runner) { long now = DateTimeOffset.UtcNow.ToUnixTimeMilliseconds(); + string adapterDir = jobConfig?["adapter_dir"]?.ToString() ?? ""; + string outputName = jobConfig?["output_name"]?.ToString() ?? ""; + JObject progress = TrainingJobManager.GetProgress(); + string finalStatus = success ? "completed" : "failed"; + if (success && Directory.Exists(adapterDir)) + { + JObject reg = await RegisterAdapterPipeline(session, baseUrl, outputName, adapterDir, jobConfig, runner); + progress["ollama"] = reg; + if (reg["success"]?.Value() != true && reg["skipped"]?.Value() != true) + { + finalStatus = "completed_with_warnings"; + } + } + progress["status"] = finalStatus; Memory.SaveTrainJob(new JObject { ["id"] = jobId, - ["status"] = success ? "completed" : "failed", + ["status"] = finalStatus, ["finished_at"] = now, - ["progress"] = TrainingJobManager.GetProgress(), + ["progress"] = progress, }); - if (success && Directory.Exists(adapterDir)) - { - try - { - await RegisterAdapterInOllama(session, baseUrl, outputName, adapterDir, ggufScript); - } - catch (Exception ex) - { - Logs.Debug($"RegisterAdapter: {ex.Message}"); - } - } if (!string.IsNullOrWhiteSpace(chatModel)) { await AssistentWarmLlm(session, baseUrl, chatModel); @@ -203,19 +256,115 @@ public partial class SwarmAssistentExtension TrainingJobManager.ClearRunning(); } - async Task RegisterAdapterInOllama(Session session, string baseUrl, string outputName, string adapterDir, string ggufScript) + async Task RegisterAdapterPipeline(Session session, string baseUrl, string outputName, string adapterDir, JObject jobConfig, JObject runner) { - string adapterFile = Directory.GetFiles(adapterDir, "*.gguf").FirstOrDefault() - ?? Directory.GetFiles(adapterDir, "adapter_model.safetensors").FirstOrDefault(); - if (string.IsNullOrWhiteSpace(adapterFile)) + string safetensors = Directory.GetFiles(adapterDir, "adapter_model.safetensors").FirstOrDefault(); + if (string.IsNullOrWhiteSpace(safetensors)) { - return; + return new JObject { ["success"] = false, ["error"] = "adapter_model.safetensors not found" }; } + string ggufPath = Directory.GetFiles(adapterDir, "*.gguf").FirstOrDefault(); + string ggufScript = runner?["gguf_script"]?.ToString()?.Trim(); + string ggufBase = jobConfig?["gguf_base_path"]?.ToString()?.Trim() ?? runner?["gguf_base_path"]?.ToString()?.Trim(); + string python = runner?["python"]?.ToString()?.Trim(); + if (string.IsNullOrWhiteSpace(python)) + { + python = "python"; + } + if (string.IsNullOrWhiteSpace(ggufPath) && !string.IsNullOrWhiteSpace(ggufScript) && File.Exists(ggufScript)) + { + if (string.IsNullOrWhiteSpace(ggufBase) || !File.Exists(ggufBase)) + { + return new JObject + { + ["success"] = false, + ["skipped"] = true, + ["error"] = "gguf_base_path не задан или файл не найден — адаптер сохранён как safetensors", + ["adapter_dir"] = adapterDir, + }; + } + ggufPath = Path.Combine(adapterDir, "adapter.gguf"); + string ggufCmd = runner?["gguf_cmd"]?.ToString()?.Trim(); + if (string.IsNullOrWhiteSpace(ggufCmd)) + { + ggufCmd = "\"{python}\" \"{script}\" \"{base}\" \"{lora}\" \"{out}\""; + } + string cmd = ggufCmd + .Replace("{python}", python, StringComparison.OrdinalIgnoreCase) + .Replace("{script}", ggufScript, StringComparison.OrdinalIgnoreCase) + .Replace("{base}", ggufBase, StringComparison.OrdinalIgnoreCase) + .Replace("{lora}", adapterDir, StringComparison.OrdinalIgnoreCase) + .Replace("{out}", ggufPath, StringComparison.OrdinalIgnoreCase); + int code = await RunShellCommandAsync(cmd, adapterDir); + if (code != 0 || !File.Exists(ggufPath)) + { + return new JObject + { + ["success"] = false, + ["error"] = $"GGUF convert failed exit={code}", + ["adapter_dir"] = adapterDir, + }; + } + } + if (string.IsNullOrWhiteSpace(ggufPath) || !File.Exists(ggufPath)) + { + return new JObject + { + ["success"] = false, + ["skipped"] = true, + ["note"] = "Настрой convert_lora_to_gguf.py и gguf_base_path для регистрации в Ollama", + ["adapter_dir"] = adapterDir, + }; + } + string ollamaBase = jobConfig?["ollama_base"]?.ToString()?.Trim(); + if (string.IsNullOrWhiteSpace(ollamaBase)) + { + return new JObject + { + ["success"] = false, + ["skipped"] = true, + ["error"] = "ollama_base не задан — укажи базовую Ollama-модель на форме QLoRA", + ["adapter_dir"] = adapterDir, + ["gguf"] = ggufPath, + }; + } + return await RegisterAdapterInOllama(baseUrl, outputName, ollamaBase, ggufPath); + } + + static async Task RunShellCommandAsync(string commandLine, string workDir) + { + try + { + ProcessStartInfo psi = new() + { + FileName = "cmd.exe", + Arguments = $"/c {commandLine}", + UseShellExecute = false, + RedirectStandardOutput = true, + RedirectStandardError = true, + CreateNoWindow = true, + WorkingDirectory = workDir ?? Environment.CurrentDirectory, + }; + using Process proc = Process.Start(psi); + if (proc is null) + { + return -1; + } + await proc.WaitForExitAsync(); + return proc.ExitCode; + } + catch (Exception ex) + { + Logs.Debug($"RunShellCommand: {ex.Message}"); + return -1; + } + } + + async Task 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 @@
+ @@ -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));