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 <cursoragent@cursor.com>
This commit is contained in:
Leonid Pershin
2026-08-22 14:45:07 +03:00
co-authored by Cursor
parent 1a03c3178f
commit 474a674a35
12 changed files with 1002 additions and 133 deletions
+116 -7
View File
@@ -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) => `<option value="${escapeHtml(c)}">${escapeHtml(c)}</option>`).join("");
if (cols.includes("tags")) userSel.value = "tags";
else if (cols.includes("title")) userSel.value = "title";
}
if (asstSel) {
asstSel.innerHTML = cols.map((c) => `<option value="${escapeHtml(c)}">${escapeHtml(c)}</option>`).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 = '<option value="">\u2014</option>';
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));
+15
View File
@@ -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;
+149 -4
View File
@@ -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<string> 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<bool> 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<string> 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();
+10 -1
View File
@@ -296,6 +296,7 @@ public partial class SwarmAssistentExtension
{
List<JObject> 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,
};
}
+244 -56
View File
@@ -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<int?>() ?? 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<int?>() ?? 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<bool?>() != true && reg["skipped"]?.Value<bool?>() != 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<JObject> 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<int> 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<JObject> 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,12 +498,29 @@ 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;
}
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
@@ -347,13 +535,11 @@ sealed class TrainingJobManager
// 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);
}
}
+12
View File
@@ -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
+3 -3
View File
@@ -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)
+23 -7
View File
@@ -222,6 +222,17 @@
<div class="sa-hf-status" id="sa_hf_status" role="status"></div>
<div class="sa-hf-list" id="sa_hf_list"></div>
<div class="sa-hf-preview" id="sa_hf_preview" hidden></div>
<div class="sa-hf-mapping-row" id="sa_hf_mapping_row" hidden>
<span class="sa-settings-hint">Маппинг колонок:</span>
<label>Preset
<select id="sa_hf_mapping_preset" class="sa-select">
<option value="">— вручную —</option>
<option value="fiction_tags_text">Fiction: title/tags → text</option>
</select>
</label>
<label>User col <select id="sa_hf_user_col" class="sa-select"></select></label>
<label>Assistant col <select id="sa_hf_asst_col" class="sa-select"></select></label>
</div>
<div class="sa-hf-import-row" id="sa_hf_import_row" hidden>
<label>Лимит строк <input type="number" id="sa_hf_import_limit" min="1" max="5000" value="200" /></label>
<label>Соотношение внешних:своих <input type="number" id="sa_hf_mix_ratio" min="0" max="20" step="0.5" value="3" title="Сколько внешних примеров на один свой" /></label>
@@ -253,17 +264,22 @@
</div>
<div class="sa-train-form" id="sa_train_form_qlora" hidden>
<p class="sa-settings-hint">Base model — HF id (safetensors). Тренер скачает веса сам (нужен HF_TOKEN в User Settings).</p>
<label>HF base model <input type="text" id="sa_qlora_base" placeholder="meta-llama/Llama-3.2-3B-Instruct" /></label>
<label>Имя адаптера <input type="text" id="sa_qlora_name" placeholder="assistent-lora-v1" /></label>
<label>HF base model <input type="text" id="sa_qlora_base" placeholder="Qwen/Qwen2.5-7B-Instruct" /></label>
<label>Ollama base (FROM для ADAPTER)
<select id="sa_qlora_ollama_base" class="sa-select"><option value=""></option></select>
</label>
<label>Имя модели в Ollama <input type="text" id="sa_qlora_name" placeholder="my-lora:v1" /></label>
<div class="sa-settings-row sa-knob-row">
<label>rank <input type="number" id="sa_qlora_rank" min="4" max="128" value="16" /></label>
<label>alpha <input type="number" id="sa_qlora_alpha" min="4" max="256" value="32" /></label>
<label>LR <input type="number" id="sa_qlora_lr" min="0.000001" max="0.01" step="0.00001" value="0.0002" /></label>
<label>epochs <input type="number" id="sa_qlora_epochs" min="1" max="20" value="3" /></label>
<label>seq_len <input type="number" id="sa_qlora_seq" min="512" max="8192" step="256" value="2048" /></label>
<label>max_samples <input type="number" id="sa_qlora_max_samples" min="0" max="50000" value="0" title="0 = все строки" /></label>
</div>
<label class="sa-check"><input type="checkbox" id="sa_qlora_4bit" checked /> 4-bit QLoRA</label>
<label>HF датасет (опционально) <input type="text" id="sa_qlora_hf_dataset" placeholder="owner/name — только для раннера" /></label>
<label>HF датасет (опционально, без локального датасета) <input type="text" id="sa_qlora_hf_dataset" placeholder="krplt/ru-fictext-nsfw" /></label>
<p class="sa-settings-hint">После обучения: safetensors → GGUF (convert script) → ollama create. Настрой пути в Настройки → Модели.</p>
<button type="button" class="basic-button sa-primary" id="sa_btn_qlora_start">Запустить QLoRA</button>
<button type="button" class="basic-button" id="sa_btn_qlora_cancel" hidden>Отменить</button>
</div>
@@ -323,15 +339,15 @@
<label>Python <input type="text" id="sa_runner_python" placeholder="python или полный путь" /></label>
<label>Тип тренера
<select id="sa_runner_kind" class="sa-select">
<option value="">— не настроен —</option>
<option value="llama-factory">LLaMA-Factory</option>
<option value="unsloth">Unsloth</option>
<option value="builtin" selected>Встроенный (train_qlora.py)</option>
<option value="custom">Custom command</option>
</select>
</label>
<label>GGUF base model (.gguf) <input type="text" id="sa_runner_gguf_base" placeholder="C:/models/base.gguf" title="Базовая GGUF для convert_lora_to_gguf" /></label>
<label>Рабочая директория <input type="text" id="sa_runner_workdir" placeholder="Assistent/training/runner" /></label>
<label>Custom command <input type="text" id="sa_runner_cmd" placeholder="{python} train.py --config {config}" /></label>
<label>convert_lora_to_gguf.py <input type="text" id="sa_runner_gguf_script" placeholder="путь к convert_lora_to_gguf.py" /></label>
<label>convert_lora_to_gguf.py <input type="text" id="sa_runner_gguf_script" placeholder="C:/llama.cpp/convert_lora_to_gguf.py" /></label>
<label>GGUF cmd template <input type="text" id="sa_runner_gguf_cmd" placeholder='"{python}" "{script}" "{base}" "{lora}" "{out}"' /></label>
<div class="sa-settings-row">
<button type="button" class="basic-button" id="sa_btn_refresh_models">Обновить модели</button>
<button type="button" class="basic-button" id="sa_btn_refresh_inventory">Обновить inventory</button>
+94
View File
@@ -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
+9
View File
@@ -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
+197 -37
View File
@@ -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
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)
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)
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)
+118 -6
View File
@@ -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) => `<option value="${escapeHtml(c)}">${escapeHtml(c)}</option>`).join('');
if (cols.includes('tags')) userSel.value = 'tags';
else if (cols.includes('title')) userSel.value = 'title';
}
if (asstSel) {
asstSel.innerHTML = cols.map((c) => `<option value="${escapeHtml(c)}">${escapeHtml(c)}</option>`).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 = '<option value="">—</option>';
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));