using System; using System.Collections.Generic; using System.Diagnostics; using System.IO; using System.Linq; using System.Text; using System.Text.RegularExpressions; using System.Threading; using System.Threading.Tasks; using Newtonsoft.Json.Linq; using SwarmUI.Accounts; using SwarmUI.Utils; using SwarmUI.WebAPI; namespace Mrleo1nid.SwarmAssistent; /// QLoRA training job runner with VRAM lock and progress streaming. public partial class SwarmAssistentExtension { static readonly TrainingJobManager TrainingJobManager = new(); public async Task AssistentStartTrainJob(Session session, JObject raw) { if (raw is null) { return new JObject { ["error"] = "body required" }; } if (TrainingJobManager.IsRunning) { return new JObject { ["error"] = "Тренировка уже идёт" }; } string hfBase = raw["base_model"]?.ToString()?.Trim(); string outputName = raw["output_name"]?.ToString()?.Trim(); if (string.IsNullOrWhiteSpace(hfBase) || string.IsNullOrWhiteSpace(outputName)) { return new JObject { ["error"] = "base_model and output_name required" }; } if (!string.IsNullOrWhiteSpace(raw["hf_dataset"]?.ToString())) { string dsId = NormalizeHfDatasetId(raw["hf_dataset"]?.ToString()); if (dsId is null) { return new JObject { ["error"] = "invalid hf_dataset id" }; } JObject check = await CheckHfDatasetInternal(session, dsId, useCache: true); if (check["gate"]?.ToString() == "rejected") { return new JObject { ["error"] = check["reason"]?.ToString() ?? "hf dataset rejected" }; } 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-раннер не настроен (Настройки → Модели)" }; } string baseUrl = NormalizeBaseUrl(raw["base_url"]?.ToString()); string chatModel = raw["chat_model"]?.ToString()?.Trim(); if (!string.IsNullOrWhiteSpace(chatModel)) { await AssistentParkLlm(session, baseUrl, chatModel); } JObject export = await AssistentExportDataset(session, "approved", "jsonl"); string datasetPath = export["path"]?.ToString(); if (string.IsNullOrWhiteSpace(datasetPath) || !File.Exists(datasetPath)) { return new JObject { ["error"] = "Нет одобренных примеров для тренировки" }; } long now = DateTimeOffset.UtcNow.ToUnixTimeMilliseconds(); string jobId = $"tj_{now}"; string jobDir = Path.Combine(TrainingRoot(), "jobs", jobId); Directory.CreateDirectory(jobDir); string configPath = Path.Combine(jobDir, "config.json"); string logPath = Path.Combine(jobDir, "log.txt"); JObject jobConfig = new() { ["base_model"] = hfBase, ["output_name"] = outputName, ["dataset_path"] = datasetPath, ["hf_dataset"] = raw["hf_dataset"], ["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), }; await File.WriteAllTextAsync(configPath, jobConfig.ToString(Newtonsoft.Json.Formatting.Indented), Encoding.UTF8); Memory.SaveTrainJob(new JObject { ["id"] = jobId, ["kind"] = "qlora", ["status"] = "running", ["config"] = jobConfig, ["base_model"] = hfBase, ["output_name"] = outputName, ["log_path"] = logPath, ["created_at"] = now, }); string cmdLine = BuildRunnerCommand(runner, configPath, logPath, jobDir); bool started = TrainingJobManager.Start(this, session, jobId, cmdLine, logPath, baseUrl, chatModel, GetHfToken(session)); if (!started) { Memory.SaveTrainJob(new JObject { ["id"] = jobId, ["status"] = "failed", ["progress"] = new JObject { ["error"] = "process start failed" } }); return new JObject { ["error"] = "Не удалось запустить процесс тренировки" }; } return new JObject { ["success"] = true, ["job_id"] = jobId, ["log_path"] = logPath }; } static string BuildRunnerCommand(JObject runner, string configPath, string logPath, string workDir) { string python = runner["python"]?.ToString()?.Trim(); if (string.IsNullOrWhiteSpace(python)) { python = "python"; } string kind = runner["kind"]?.ToString()?.Trim() ?? "custom"; string custom = runner["cmd"]?.ToString()?.Trim(); string scriptPath = Path.Combine(FilePath, "scripts", "train_qlora.py"); if (kind == "custom" && !string.IsNullOrWhiteSpace(custom)) { return custom .Replace("{python}", python, StringComparison.OrdinalIgnoreCase) .Replace("{config}", configPath, StringComparison.OrdinalIgnoreCase) .Replace("{log}", logPath, StringComparison.OrdinalIgnoreCase) .Replace("{workdir}", workDir, StringComparison.OrdinalIgnoreCase); } return $"\"{python}\" \"{scriptPath}\" --config \"{configPath}\" --log \"{logPath}\""; } public async Task AssistentCancelTrainJob(Session session, string id = null) { await Task.CompletedTask; TrainingJobManager.Cancel(); string jobId = id ?? TrainingJobManager.CurrentJobId; if (!string.IsNullOrWhiteSpace(jobId)) { Memory.SaveTrainJob(new JObject { ["id"] = jobId, ["status"] = "cancelled", ["finished_at"] = DateTimeOffset.UtcNow.ToUnixTimeMilliseconds(), }); } return new JObject { ["success"] = true, ["cancelled"] = true }; } public async Task AssistentTrainWS(Session session, WebSocket ws, JObject raw) { await Task.CompletedTask; try { while (TrainingJobManager.IsRunning && ws.State == System.Net.WebSockets.WebSocketState.Open) { JObject progress = TrainingJobManager.GetProgress(); string msg = progress.ToString(Newtonsoft.Json.Formatting.None); await ws.SendAsync(Encoding.UTF8.GetBytes(msg), System.Net.WebSockets.WebSocketMessageType.Text, true, CancellationToken.None); await Task.Delay(800); } JObject final = TrainingJobManager.GetProgress(); final["done"] = true; await ws.SendAsync(Encoding.UTF8.GetBytes(final.ToString(Newtonsoft.Json.Formatting.None)), System.Net.WebSockets.WebSocketMessageType.Text, true, CancellationToken.None); } catch (Exception ex) { Logs.Debug($"AssistentTrainWS: {ex.Message}"); } 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) { long now = DateTimeOffset.UtcNow.ToUnixTimeMilliseconds(); Memory.SaveTrainJob(new JObject { ["id"] = jobId, ["status"] = success ? "completed" : "failed", ["finished_at"] = now, ["progress"] = TrainingJobManager.GetProgress(), }); 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); } TrainingJobManager.ClearRunning(); } async Task RegisterAdapterInOllama(Session session, string baseUrl, string outputName, string adapterDir, string ggufScript) { string adapterFile = Directory.GetFiles(adapterDir, "*.gguf").FirstOrDefault() ?? Directory.GetFiles(adapterDir, "adapter_model.safetensors").FirstOrDefault(); if (string.IsNullOrWhiteSpace(adapterFile)) { return; } 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("\\", "/")}"); JObject payload = new() { ["name"] = outputName, ["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($"{NormalizeBaseUrl(baseUrl)}/api/create", content); _ = await resp.Content.ReadAsStringAsync(); } } 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); Process _process; readonly object _lock = new(); JObject _progress = new() { ["status"] = "idle" }; string _logPath; SwarmAssistentExtension _ext; Session _session; string _jobId; string _baseUrl; string _chatModel; public bool IsRunning { get; private set; } public string CurrentJobId => _jobId; public bool Start(SwarmAssistentExtension ext, Session session, string jobId, string commandLine, string logPath, string baseUrl, string chatModel, string hfToken) { lock (_lock) { if (IsRunning) { return false; } _ext = ext; _session = session; _jobId = jobId; _logPath = logPath; _baseUrl = baseUrl; _chatModel = chatModel; _progress = new JObject { ["status"] = "running", ["step"] = 0, ["loss"] = null, ["log"] = "" }; try { ProcessStartInfo psi = new() { FileName = "cmd.exe", Arguments = $"/c {commandLine}", UseShellExecute = false, RedirectStandardOutput = true, RedirectStandardError = true, CreateNoWindow = true, WorkingDirectory = Path.GetDirectoryName(logPath) ?? Environment.CurrentDirectory, }; if (!string.IsNullOrWhiteSpace(hfToken)) { psi.Environment["HF_TOKEN"] = hfToken; } _process = new Process { StartInfo = psi, EnableRaisingEvents = true }; _process.OutputDataReceived += (_, e) => AppendLog(e.Data); _process.ErrorDataReceived += (_, e) => AppendLog(e.Data); _process.Exited += async (_, _) => await OnExited(); _process.Start(); _process.BeginOutputReadLine(); _process.BeginErrorReadLine(); IsRunning = true; return true; } catch (Exception ex) { _progress["error"] = ex.Message; IsRunning = false; return false; } } } void AppendLog(string line) { if (string.IsNullOrWhiteSpace(line)) { return; } lock (_lock) { try { File.AppendAllText(_logPath, line + Environment.NewLine); } catch { // ignore } string prev = _progress["log"]?.ToString() ?? ""; string combined = (prev + line + "\n"); if (combined.Length > 12000) { combined = combined[^12000..]; } _progress["log"] = combined; Match lossM = LossRe.Match(line); if (lossM.Success) { _progress["loss"] = lossM.Groups[1].Value; } 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); _progress["percent"] = total > 0 ? (int)(100.0 * step / total) : 0; } try { _ext?.Memory?.SaveTrainJob(new JObject { ["id"] = _jobId, ["status"] = "running", ["progress"] = _progress, }); } catch { // ignore } } } async Task OnExited() { bool ok = false; string adapterDir = ""; string outputName = ""; lock (_lock) { ok = _process?.ExitCode == 0; IsRunning = false; _progress["status"] = ok ? "completed" : "failed"; _progress["exit_code"] = _process?.ExitCode; } if (_ext != null) { 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() ?? ""; } catch { // ignore } JObject runner = _ext.Config.LoadTrainingRunner(); await _ext.FinishTrainJobAsync(_jobId, ok, _logPath, _session, _baseUrl, _chatModel, adapterDir, outputName, runner["gguf_script"]?.ToString()); } } public void Cancel() { lock (_lock) { try { if (_process != null && !_process.HasExited) { _process.Kill(entireProcessTree: true); } } catch { // ignore } IsRunning = false; _progress["status"] = "cancelled"; } } public JObject GetProgress() { lock (_lock) { return (JObject)_progress.DeepClone(); } } public void ClearRunning() { lock (_lock) { IsRunning = false; } } }