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>
453 lines
17 KiB
C#
453 lines
17 KiB
C#
using System;
|
|
using System.Collections.Generic;
|
|
using System.IO;
|
|
using System.Linq;
|
|
using System.Text;
|
|
using System.Threading.Tasks;
|
|
using Newtonsoft.Json.Linq;
|
|
using SwarmUI.Accounts;
|
|
using SwarmUI.Utils;
|
|
|
|
namespace Mrleo1nid.SwarmAssistent;
|
|
|
|
/// <summary>Training samples, dataset import/export, Ollama Modelfile builder.</summary>
|
|
public partial class SwarmAssistentExtension
|
|
{
|
|
static string TrainingRoot()
|
|
{
|
|
string root = Path.Combine(DataRoot(), "Assistent", "training");
|
|
Directory.CreateDirectory(root);
|
|
Directory.CreateDirectory(Path.Combine(root, "datasets"));
|
|
Directory.CreateDirectory(Path.Combine(root, "jobs"));
|
|
Directory.CreateDirectory(Path.Combine(root, "adapters"));
|
|
return root;
|
|
}
|
|
|
|
public async Task<JObject> AssistentListTrainSamples(Session session, string status = null, string persona = null, int limit = 200)
|
|
{
|
|
await Task.CompletedTask;
|
|
try
|
|
{
|
|
List<JObject> list = Memory.ListTrainSamples(status, persona, null, limit);
|
|
return new JObject
|
|
{
|
|
["success"] = true,
|
|
["samples"] = new JArray(list),
|
|
["approved"] = Memory.CountTrainSamples("approved"),
|
|
["total"] = Memory.CountTrainSamples(null),
|
|
};
|
|
}
|
|
catch (Exception ex)
|
|
{
|
|
return new JObject { ["error"] = ex.Message };
|
|
}
|
|
}
|
|
|
|
public async Task<JObject> AssistentUpsertTrainSample(Session session, JObject raw)
|
|
{
|
|
await Task.CompletedTask;
|
|
if (raw is null)
|
|
{
|
|
return new JObject { ["error"] = "body required" };
|
|
}
|
|
try
|
|
{
|
|
JObject saved = Memory.UpsertTrainSample(raw);
|
|
JObject full = Memory.GetTrainSample(saved["id"]?.ToString()) ?? raw;
|
|
await TryAutoLinkTrainSample(session, full, raw);
|
|
return new JObject { ["success"] = true, ["sample"] = full };
|
|
}
|
|
catch (Exception ex)
|
|
{
|
|
return new JObject { ["error"] = ex.Message };
|
|
}
|
|
}
|
|
|
|
public async Task<JObject> AssistentDeleteTrainSample(Session session, string id)
|
|
{
|
|
await Task.CompletedTask;
|
|
if (string.IsNullOrWhiteSpace(id))
|
|
{
|
|
return new JObject { ["error"] = "id required" };
|
|
}
|
|
JObject sample = Memory.GetTrainSample(id.Trim());
|
|
bool ok = Memory.DeleteTrainSample(id.Trim());
|
|
if (ok && sample is not null)
|
|
{
|
|
Memory.UnlinkTrainSampleFromAgent(sample);
|
|
}
|
|
return new JObject { ["success"] = true, ["deleted"] = ok };
|
|
}
|
|
|
|
public async Task<JObject> AssistentBuildDatasetFromChats(Session session, bool approved_only = false)
|
|
{
|
|
await Task.CompletedTask;
|
|
try
|
|
{
|
|
List<JObject> chats = Memory.ListChats(withMessages: true, limit: AssistentMemory.MaxChatsStored);
|
|
int added = 0;
|
|
foreach (JObject chat in chats)
|
|
{
|
|
JArray messages = chat["messages"] as JArray ?? [];
|
|
for (int i = 0; i < messages.Count - 1; i++)
|
|
{
|
|
if (messages[i] is not JObject u || messages[i + 1] is not JObject a)
|
|
{
|
|
continue;
|
|
}
|
|
if (!string.Equals(u["role"]?.ToString(), "user", StringComparison.OrdinalIgnoreCase))
|
|
{
|
|
continue;
|
|
}
|
|
if (!string.Equals(a["role"]?.ToString(), "assistant", StringComparison.OrdinalIgnoreCase))
|
|
{
|
|
continue;
|
|
}
|
|
Memory.UpsertTrainSample(new JObject
|
|
{
|
|
["source"] = "chat",
|
|
["chat_id"] = chat["id"],
|
|
["persona"] = a["persona"] ?? u["persona"],
|
|
["pack"] = a["pack"] ?? u["pack"],
|
|
["status"] = approved_only ? "approved" : "draft",
|
|
["messages"] = new JArray { u.DeepClone(), a.DeepClone() },
|
|
});
|
|
added++;
|
|
}
|
|
}
|
|
return new JObject { ["success"] = true, ["added"] = added };
|
|
}
|
|
catch (Exception ex)
|
|
{
|
|
return new JObject { ["error"] = ex.Message };
|
|
}
|
|
}
|
|
|
|
public async Task<JObject> AssistentImportDataset(Session session, JObject raw)
|
|
{
|
|
await Task.CompletedTask;
|
|
string content = raw?["content"]?.ToString();
|
|
string format = raw?["format"]?.ToString() ?? "auto";
|
|
if (string.IsNullOrWhiteSpace(content))
|
|
{
|
|
return new JObject { ["error"] = "content required" };
|
|
}
|
|
try
|
|
{
|
|
int imported = 0;
|
|
string fmt = (format ?? "auto").Trim().ToLowerInvariant();
|
|
List<JObject> records = ParseDatasetContent(content, fmt);
|
|
foreach (JObject rec in records)
|
|
{
|
|
JArray messages = rec["messages"] as JArray;
|
|
if (messages is null || messages.Count == 0)
|
|
{
|
|
continue;
|
|
}
|
|
Memory.UpsertTrainSample(new JObject
|
|
{
|
|
["source"] = "import",
|
|
["messages"] = messages,
|
|
["status"] = "draft",
|
|
});
|
|
imported++;
|
|
}
|
|
return new JObject { ["success"] = true, ["imported"] = imported };
|
|
}
|
|
catch (Exception ex)
|
|
{
|
|
return new JObject { ["error"] = ex.Message };
|
|
}
|
|
}
|
|
|
|
static List<JObject> ParseDatasetContent(string content, string format)
|
|
{
|
|
List<JObject> list = [];
|
|
string trimmed = content.Trim();
|
|
if (trimmed.StartsWith('['))
|
|
{
|
|
JArray arr = JArray.Parse(trimmed);
|
|
foreach (JToken t in arr)
|
|
{
|
|
if (t is JObject o)
|
|
{
|
|
JArray msgs = ExtractMessagesFromRecord(o);
|
|
if (msgs != null)
|
|
{
|
|
list.Add(new JObject { ["messages"] = msgs });
|
|
}
|
|
}
|
|
}
|
|
return list;
|
|
}
|
|
if (format == "csv" || LooksLikeCsv(trimmed))
|
|
{
|
|
return ParseCsvDataset(trimmed);
|
|
}
|
|
foreach (string line in trimmed.Split('\n'))
|
|
{
|
|
string ln = line.Trim();
|
|
if (string.IsNullOrWhiteSpace(ln))
|
|
{
|
|
continue;
|
|
}
|
|
try
|
|
{
|
|
JObject o = JObject.Parse(ln);
|
|
JArray msgs = ExtractMessagesFromRecord(o);
|
|
if (msgs != null)
|
|
{
|
|
list.Add(new JObject { ["messages"] = msgs });
|
|
}
|
|
}
|
|
catch
|
|
{
|
|
// skip bad line
|
|
}
|
|
}
|
|
return list;
|
|
}
|
|
|
|
static bool LooksLikeCsv(string s) => s.Contains(',') && s.Contains('\n') && !s.TrimStart().StartsWith('{');
|
|
|
|
static List<JObject> ParseCsvDataset(string csv)
|
|
{
|
|
List<JObject> list = [];
|
|
string[] lines = csv.Split('\n').Select(l => l.Trim()).Where(l => l.Length > 0).ToArray();
|
|
if (lines.Length < 2)
|
|
{
|
|
return list;
|
|
}
|
|
string[] headers = lines[0].Split(',').Select(h => h.Trim().Trim('"')).ToArray();
|
|
int promptIdx = Array.FindIndex(headers, h => h.Equals("prompt", StringComparison.OrdinalIgnoreCase) || h.Equals("question", StringComparison.OrdinalIgnoreCase) || h.Equals("instruction", StringComparison.OrdinalIgnoreCase));
|
|
int respIdx = Array.FindIndex(headers, h => h.Equals("response", StringComparison.OrdinalIgnoreCase) || h.Equals("answer", StringComparison.OrdinalIgnoreCase) || h.Equals("output", StringComparison.OrdinalIgnoreCase) || h.Equals("completion", StringComparison.OrdinalIgnoreCase));
|
|
if (promptIdx < 0 || respIdx < 0)
|
|
{
|
|
return list;
|
|
}
|
|
for (int i = 1; i < lines.Length; i++)
|
|
{
|
|
string[] cols = SplitCsvLine(lines[i]);
|
|
if (cols.Length <= Math.Max(promptIdx, respIdx))
|
|
{
|
|
continue;
|
|
}
|
|
list.Add(new JObject
|
|
{
|
|
["messages"] = new JArray
|
|
{
|
|
new JObject { ["role"] = "user", ["content"] = cols[promptIdx] },
|
|
new JObject { ["role"] = "assistant", ["content"] = cols[respIdx] },
|
|
},
|
|
});
|
|
}
|
|
return list;
|
|
}
|
|
|
|
static string[] SplitCsvLine(string line)
|
|
{
|
|
List<string> parts = [];
|
|
StringBuilder cur = new();
|
|
bool inQ = false;
|
|
foreach (char c in line)
|
|
{
|
|
if (c == '"')
|
|
{
|
|
inQ = !inQ;
|
|
continue;
|
|
}
|
|
if (c == ',' && !inQ)
|
|
{
|
|
parts.Add(cur.ToString().Trim());
|
|
cur.Clear();
|
|
continue;
|
|
}
|
|
cur.Append(c);
|
|
}
|
|
parts.Add(cur.ToString().Trim());
|
|
return parts.ToArray();
|
|
}
|
|
|
|
static JArray ExtractMessagesFromRecord(JObject o)
|
|
{
|
|
if (o["messages"] is JArray msgs)
|
|
{
|
|
return NormalizeMessagesArray(msgs);
|
|
}
|
|
if (o["conversations"] is JArray conv)
|
|
{
|
|
return ConvertHfRowToMessages(new JObject { ["conversations"] = conv }, new JObject { ["kind"] = "conversations" }, null);
|
|
}
|
|
if (o["instruction"] != null && o["output"] != null)
|
|
{
|
|
return ConvertHfRowToMessages(o, new JObject { ["kind"] = "alpaca" }, null);
|
|
}
|
|
if (o["prompt"] != null && (o["response"] != null || o["completion"] != null))
|
|
{
|
|
return ConvertHfRowToMessages(o, new JObject { ["kind"] = "prompt_response", ["response_col"] = o["response"] != null ? "response" : "completion" }, null);
|
|
}
|
|
return null;
|
|
}
|
|
|
|
public async Task<JObject> AssistentExportDataset(Session session, string status = "approved", string format = "jsonl")
|
|
{
|
|
await Task.CompletedTask;
|
|
try
|
|
{
|
|
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 ?? [];
|
|
if (messages.Count == 0)
|
|
{
|
|
continue;
|
|
}
|
|
exported++;
|
|
if (string.Equals(format, "sharegpt", StringComparison.OrdinalIgnoreCase))
|
|
{
|
|
JArray conv = [];
|
|
foreach (JToken m in messages)
|
|
{
|
|
if (m is not JObject mo)
|
|
{
|
|
continue;
|
|
}
|
|
string role = mo["role"]?.ToString() ?? "user";
|
|
conv.Add(new JObject
|
|
{
|
|
["from"] = role == "assistant" ? "gpt" : "human",
|
|
["value"] = mo["content"]?.ToString() ?? "",
|
|
});
|
|
}
|
|
sb.AppendLine(new JObject { ["conversations"] = conv }.ToString(Newtonsoft.Json.Formatting.None));
|
|
}
|
|
else
|
|
{
|
|
sb.AppendLine(new JObject { ["messages"] = messages }.ToString(Newtonsoft.Json.Formatting.None));
|
|
}
|
|
}
|
|
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"] = exported, ["content"] = sb.ToString() };
|
|
}
|
|
catch (Exception ex)
|
|
{
|
|
return new JObject { ["error"] = ex.Message };
|
|
}
|
|
}
|
|
|
|
public async Task<JObject> AssistentCreateOllamaModel(Session session, JObject raw)
|
|
{
|
|
if (raw is null)
|
|
{
|
|
return new JObject { ["error"] = "body required" };
|
|
}
|
|
string baseUrl = NormalizeBaseUrl(raw["base_url"]?.ToString());
|
|
string baseModel = raw["base_model"]?.ToString()?.Trim();
|
|
string name = raw["name"]?.ToString()?.Trim();
|
|
string system = raw["system"]?.ToString() ?? "";
|
|
int shots = raw["shots"]?.Value<int?>() ?? 8;
|
|
if (string.IsNullOrWhiteSpace(baseModel) || string.IsNullOrWhiteSpace(name))
|
|
{
|
|
return new JObject { ["error"] = "base_model and name required" };
|
|
}
|
|
if (string.IsNullOrWhiteSpace(system))
|
|
{
|
|
string persona = AssistentConfig.SafeId(raw["persona"]?.ToString()) ?? Config.DefaultPersonaId();
|
|
system = Config.LoadCorePrompt(persona) + "\n\n" + Config.RenderIdentityBlock(persona, includeAllShelves: true);
|
|
}
|
|
StringBuilder mf = new();
|
|
mf.AppendLine($"FROM {baseModel}");
|
|
mf.AppendLine($"SYSTEM \"\"\"{system}\"\"\"");
|
|
List<JObject> samples = Memory.ListTrainSamples("approved", null, null, Math.Clamp(shots, 0, 32));
|
|
foreach (JObject s in samples.Take(shots))
|
|
{
|
|
JArray messages = s["messages"] as JArray ?? [];
|
|
foreach (JToken m in messages)
|
|
{
|
|
if (m is not JObject mo)
|
|
{
|
|
continue;
|
|
}
|
|
string role = mo["role"]?.ToString() ?? "user";
|
|
string content = mo["content"]?.ToString() ?? "";
|
|
if (string.IsNullOrWhiteSpace(content))
|
|
{
|
|
continue;
|
|
}
|
|
mf.AppendLine($"MESSAGE {role} \"\"\"{content.Replace("\"\"\"", "\"\"\"\"\"\"\"")}\"\"\"");
|
|
}
|
|
}
|
|
if (raw["num_ctx"] != null)
|
|
{
|
|
mf.AppendLine($"PARAMETER num_ctx {raw["num_ctx"]}");
|
|
}
|
|
if (raw["temperature"] != null)
|
|
{
|
|
mf.AppendLine($"PARAMETER temperature {raw["temperature"]}");
|
|
}
|
|
string modelfilePath = Path.Combine(TrainingRoot(), "jobs", $"modelfile_{DateTimeOffset.UtcNow.ToUnixTimeMilliseconds()}.Modelfile");
|
|
Directory.CreateDirectory(Path.GetDirectoryName(modelfilePath)!);
|
|
await File.WriteAllTextAsync(modelfilePath, mf.ToString(), Encoding.UTF8);
|
|
try
|
|
{
|
|
JObject payload = new()
|
|
{
|
|
["name"] = name,
|
|
["modelfile"] = mf.ToString(),
|
|
["stream"] = false,
|
|
};
|
|
using StringContent content = new(payload.ToString(Newtonsoft.Json.Formatting.None), Encoding.UTF8, "application/json");
|
|
using HttpResponseMessage resp = await HttpClient.PostAsync($"{baseUrl}/api/create", content);
|
|
string body = await resp.Content.ReadAsStringAsync();
|
|
if (!resp.IsSuccessStatusCode)
|
|
{
|
|
return new JObject { ["error"] = $"Ollama create HTTP {(int)resp.StatusCode}: {Clip(body, 400)}", ["modelfile_path"] = modelfilePath };
|
|
}
|
|
return new JObject { ["success"] = true, ["name"] = name, ["modelfile_path"] = modelfilePath, ["ollama"] = body };
|
|
}
|
|
catch (Exception ex)
|
|
{
|
|
return new JObject { ["error"] = ex.Message, ["modelfile_path"] = modelfilePath };
|
|
}
|
|
}
|
|
|
|
public async Task<JObject> AssistentGetTrainJob(Session session, string id = null)
|
|
{
|
|
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,
|
|
};
|
|
}
|
|
|
|
public async Task<JObject> AssistentSaveRunnerSettings(Session session, JObject settings)
|
|
{
|
|
await Task.CompletedTask;
|
|
if (settings is null)
|
|
{
|
|
return new JObject { ["error"] = "settings required" };
|
|
}
|
|
Config.SaveTrainingRunner(settings);
|
|
return new JObject { ["success"] = true };
|
|
}
|
|
|
|
public async Task<JObject> AssistentGetRunnerSettings(Session session)
|
|
{
|
|
await Task.CompletedTask;
|
|
return new JObject { ["success"] = true, ["settings"] = Config.LoadTrainingRunner() };
|
|
}
|
|
}
|