From f4ebfeb757460fdfea0c30de5c14bd9e75fea5b5 Mon Sep 17 00:00:00 2001 From: Josh Lavely Date: Thu, 24 Sep 2026 21:11:12 -0400 Subject: [PATCH 01/14] feat: update refine_faces function and add unit tests for image processing --- runtime/image_worker.py | 8 +-- runtime/requirements-image.txt | 3 +- tests/python/test_image_worker.py | 84 +++++++++++++++++++++++++++++++ 3 files changed, 88 insertions(+), 7 deletions(-) create mode 100644 tests/python/test_image_worker.py diff --git a/runtime/image_worker.py b/runtime/image_worker.py index 5d33246..faf0ecb 100644 --- a/runtime/image_worker.py +++ b/runtime/image_worker.py @@ -235,8 +235,6 @@ def refine_faces( prompt: str, negative_prompt: str, strength: float, - steps: int, - guidance: float, generator: Any, ): import torch @@ -283,8 +281,8 @@ def refine_faces( negative_prompt=negative_prompt or None, image=resized, strength=strength, - num_inference_steps=max(20, steps), - guidance_scale=max(5.0, guidance - 1.5), + num_inference_steps=25, + guidance_scale=6.0, generator=generator, ).images[0] refined = refined.resize( @@ -501,8 +499,6 @@ def legacy_progress_callback(step_index, _timestep, latents): prompt, negative_prompt, face_fix_strength, - steps, - guidance, generator, ) processing_suffixes.append("face") diff --git a/runtime/requirements-image.txt b/runtime/requirements-image.txt index 40be050..ae93592 100644 --- a/runtime/requirements-image.txt +++ b/runtime/requirements-image.txt @@ -1,5 +1,6 @@ # Install a CUDA-compatible PyTorch build for your driver first. -torch>=2.5,<3 +torch>=2.5,<2.6 +torchvision>=0.20,<0.21 accelerate>=1.2,<2 diffusers>=0.32,<1 transformers>=4.47,<5 diff --git a/tests/python/test_image_worker.py b/tests/python/test_image_worker.py new file mode 100644 index 0000000..b735112 --- /dev/null +++ b/tests/python/test_image_worker.py @@ -0,0 +1,84 @@ +import tempfile +import unittest +from pathlib import Path +from unittest.mock import patch + +import torch +from PIL import Image + +from runtime.image_worker import ( + expanded_face_box, + feathered_mask, + require_model_file, + upscale_image, +) + + +class FakeUpscaler: + scale = 4 + + def to(self, _device): + return self + + def eval(self): + return self + + def __call__(self, image): + return torch.nn.functional.interpolate(image, scale_factor=4, mode="nearest") + + +class FakeModelLoader: + def load_from_file(self, _model_path): + return FakeUpscaler() + + +class ImageWorkerTests(unittest.TestCase): + def test_expanded_face_box_is_square_and_clamped(self): + self.assertEqual( + expanded_face_box((0, 5, 30, 25), 100, 80), + (0, 0, 36, 36), + ) + + def test_feathered_mask_keeps_center_and_softens_edges(self): + mask = feathered_mask((100, 100), feather=20) + + self.assertEqual(mask.getpixel((50, 50)), 255) + self.assertLess(mask.getpixel((0, 0)), mask.getpixel((20, 20))) + + def test_require_model_file_checks_path_and_extension(self): + with tempfile.TemporaryDirectory() as directory: + model_path = Path(directory) / "face.pt" + model_path.write_bytes(b"fixture") + invalid_path = Path(directory) / "face.bin" + invalid_path.write_bytes(b"fixture") + + self.assertEqual( + require_model_file(model_path, "Face detector", {".pt"}), + model_path.resolve(), + ) + with self.assertRaisesRegex(ValueError, "must use one of"): + require_model_file(invalid_path, "Face detector", {".pt"}) + + def test_four_x_model_can_produce_exact_two_x_output(self): + source = Image.new("RGB", (17, 13), (32, 64, 96)) + import spandrel + + with ( + patch.object(spandrel, "ImageModelDescriptor", FakeUpscaler), + patch.object(spandrel, "ModelLoader", FakeModelLoader), + patch.object(torch.cuda, "is_available", return_value=False), + ): + result = upscale_image( + source, + Path("fixture.pth"), + factor=2, + tile_size=8, + tile_overlap=2, + ) + + self.assertEqual(result.size, (34, 26)) + self.assertEqual(result.getpixel((10, 10)), (32, 64, 96)) + + +if __name__ == "__main__": + unittest.main() \ No newline at end of file From 447ad632c4de2a3e8921a40983542f1a86621b62 Mon Sep 17 00:00:00 2001 From: Josh Lavely Date: Thu, 24 Sep 2026 23:33:14 -0400 Subject: [PATCH 02/14] feat: add NSFW segmentation support with model management - Introduced NSFW segmentation functionality in the LocalJobManager. - Added model validation for NSFW segmenter files during job execution. - Implemented cleanup for stale preview images without affecting final outputs. - Enhanced the UI to allow users to install and configure NSFW segmentation models. - Updated tests to cover new NSFW segmentation features and model requirements. - Modified the image worker to handle NSFW segmentation and save detected masks. - Improved error handling and user feedback for model installation processes. --- README.md | 3 +- docs/THIRD_PARTY_ASSETS.md | 14 ++ electron/content-security-policy.test.ts | 26 +++ electron/enhancement-models.test.ts | 121 +++++++++++ electron/enhancement-models.ts | 199 +++++++++++++++++++ electron/jobs.test.ts | 55 ++++- electron/jobs.ts | 77 ++++++- electron/main.ts | 50 ++++- electron/preload.ts | 18 ++ index.html | 2 +- runtime/image_worker.py | 79 ++++++++ src/App.css | 28 +++ src/App.tsx | 82 ++++++++ src/features/operations/OperationalViews.tsx | 36 +++- src/features/studio/StudioView.tsx | 120 +++++++++-- src/lib/forge-api.ts | 13 ++ src/state/workspace.test.ts | 2 + src/state/workspace.ts | 5 + src/types.ts | 23 +++ tests/fixtures/local-job-worker.mjs | 2 +- tests/python/test_image_worker.py | 50 +++++ 21 files changed, 974 insertions(+), 31 deletions(-) create mode 100644 electron/content-security-policy.test.ts create mode 100644 electron/enhancement-models.test.ts create mode 100644 electron/enhancement-models.ts diff --git a/README.md b/README.md index bc17ec7..b95f4c7 100644 --- a/README.md +++ b/README.md @@ -28,6 +28,7 @@ Local Forge is a fast, local-first desktop workbench for open models. Chat with - Persistent sessions, image attachments, search, pinning, and deletion - Local Diffusers and `.safetensors` / `.ckpt` model registration - Cancellable image generation with live step progress and Library import +- Optional Face Fix, tiled 2x/4x UltraSharp upscaling, and saved NSFW region masks - Cancellable QLoRA training with live progress and standard PEFT output - Screenshot and reference-image imports shared by Studio and Library - Live RAM and NVIDIA telemetry, with unavailable data left unknown @@ -58,7 +59,7 @@ python -m pip install -r runtime/requirements-training.txt Select that interpreter under **Settings > Runtimes > Python runtime**. Local Forge never installs Python packages or downloads model weights automatically. -Studio executes local Diffusers directories containing `model_index.json`. Registered `.safetensors` and `.ckpt` files remain visible for inventory but must be converted to Diffusers format before generation. Completed PNGs are added to Library and stored under Local Forge's user-data `outputs/images` directory. +Studio executes local Diffusers directories containing `model_index.json`. Registered `.safetensors` and `.ckpt` files remain visible for inventory but must be converted to Diffusers format before generation. Completed PNGs are added to Library and stored under Local Forge's user-data `outputs/images` directory. Optional NSFW segmentation writes one binary mask per detected region beside the completed image. Tune requires a local Hugging Face Transformers model directory containing `config.json`; Ollama tags and GGUF files are inference artifacts and are not offered as trainable base models. Datasets may be JSON, JSONL, or CSV and should contain either a `text` field or chat `messages`. Completed adapters use standard PEFT format under `outputs/adapters`. diff --git a/docs/THIRD_PARTY_ASSETS.md b/docs/THIRD_PARTY_ASSETS.md index 22e9c4a..54bd192 100644 --- a/docs/THIRD_PARTY_ASSETS.md +++ b/docs/THIRD_PARTY_ASSETS.md @@ -17,6 +17,20 @@ The files are `1200 x 1800`, `1200 x 800`, `1200 x 800`, and `1200 x 1800` respe Images imported by a user are copied into that user's private Local Forge data directory. They are never included in the source tree or project releases. +## Optional enhancement models + +Local Forge does not bundle enhancement weights. When a user chooses Install in Studio, the app downloads an immutable, checksum-verified revision to that user's private Local Forge data directory: + +| Purpose | Repository | File | SHA-256 | +| --- | --- | --- | --- | +| 4x image upscaling | [lokCX/4x-Ultrasharp](https://huggingface.co/lokCX/4x-Ultrasharp) | `4x-UltraSharp.pth` | `a5812231fc936b42af08a5edba784195495d303d5b3248c24489ef0c4021fe01` | +| Face detection | [Bingsu/adetailer](https://huggingface.co/Bingsu/adetailer) | `face_yolov8n.pt` | `70b640f8f60b1cf0dcc72f30caf3da9495eb2fb6509da48c53374ad6806e6a9c` | +| Breast segmentation | [NSFW-API/NSFW_Segmentation](https://huggingface.co/NSFW-API/NSFW_Segmentation) | `nsfw-seg-breast-x.pt` | `5ad882ddaf149873be131943b373da9f15a0603c91508e2291ece83d729c8ecc` | +| Penis segmentation | [NSFW-API/NSFW_Segmentation](https://huggingface.co/NSFW-API/NSFW_Segmentation) | `nsfw-seg-penis-x.pt` | `49d9fc8ee67d3bdee44e46bf75aeb76058f5cd074bac027b55ff071b63a32e21` | +| Vagina segmentation | [NSFW-API/NSFW_Segmentation](https://huggingface.co/NSFW-API/NSFW_Segmentation) | `nsfw-seg-vagina-x.pt` | `f8349ae348f5cf041809c38bf38d2c80a0f3dccdb1375b97971358cde742197f` | + +These files remain subject to the terms published by their respective authors and are not covered by Local Forge's MIT code license. + ## Fonts and icons - Manrope and Space Grotesk are bundled through Fontsource packages and distributed under the SIL Open Font License 1.1. diff --git a/electron/content-security-policy.test.ts b/electron/content-security-policy.test.ts new file mode 100644 index 0000000..99a7d37 --- /dev/null +++ b/electron/content-security-policy.test.ts @@ -0,0 +1,26 @@ +import { readFile } from "node:fs/promises"; +import path from "node:path"; +import { describe, expect, it } from "vitest"; + +describe("content security policy", () => { + it("allows both Local Forge image protocols", async () => { + const html = await readFile(path.join(process.cwd(), "index.html"), "utf8"); + const page = new DOMParser().parseFromString(html, "text/html"); + const policy = page + .querySelector('meta[http-equiv="Content-Security-Policy"]') + ?.getAttribute("content"); + const imageSources = policy + ?.split(";") + .find((directive) => directive.trim().startsWith("img-src ")) + ?.trim() + .split(/\s+/) + .slice(1); + + expect(imageSources).toEqual( + expect.arrayContaining([ + "local-forge-attachment:", + "local-forge-output:", + ]), + ); + }); +}); diff --git a/electron/enhancement-models.test.ts b/electron/enhancement-models.test.ts new file mode 100644 index 0000000..86e61b4 --- /dev/null +++ b/electron/enhancement-models.test.ts @@ -0,0 +1,121 @@ +import { promises as fs } from "node:fs"; +import { createHash } from "node:crypto"; +import os from "node:os"; +import path from "node:path"; +import { afterEach, describe, expect, it } from "vitest"; +import { + discoverEnhancementModels, + downloadTrustedAsset, + type TrustedAsset, +} from "./enhancement-models"; + +const temporaryDirectories: string[] = []; + +async function temporaryDirectory(): Promise { + const directory = await fs.mkdtemp( + path.join(os.tmpdir(), "local-forge-enhancements-"), + ); + temporaryDirectories.push(directory); + return directory; +} + +afterEach(async () => { + await Promise.all( + temporaryDirectories + .splice(0) + .map((directory) => fs.rm(directory, { recursive: true, force: true })), + ); +}); + +describe("discoverEnhancementModels", () => { + it("discovers the Lavely enhancement model layout", async () => { + const root = await temporaryDirectory(); + const upscaler = path.join(root, "upscalers", "4x-UltraSharp.pth"); + const faceDetector = path.join(root, "face_detector", "face_yolov8n.pt"); + const segmenterDirectory = path.join(root, "nsfw_segmentation"); + await fs.mkdir(path.dirname(upscaler), { recursive: true }); + await fs.mkdir(path.dirname(faceDetector), { recursive: true }); + await fs.mkdir(segmenterDirectory, { recursive: true }); + await fs.writeFile(upscaler, "upscaler"); + await fs.writeFile(faceDetector, "detector"); + await Promise.all( + ["breast", "penis", "vagina"].map((region) => + fs.writeFile( + path.join(segmenterDirectory, `nsfw-seg-${region}-x.pt`), + region, + ), + ), + ); + + await expect(discoverEnhancementModels([root])).resolves.toEqual({ + upscalerModelPath: upscaler, + faceDetectorModelPath: faceDetector, + nsfwSegmenterModelPath: segmenterDirectory, + }); + }); + + it("ignores missing and incomplete model files", async () => { + const root = await temporaryDirectory(); + const upscaler = path.join(root, "upscalers", "4x-UltraSharp.pth"); + const segmenterDirectory = path.join(root, "nsfw_segmentation"); + await fs.mkdir(path.dirname(upscaler), { recursive: true }); + await fs.mkdir(segmenterDirectory, { recursive: true }); + await fs.writeFile(upscaler, ""); + await fs.writeFile( + path.join(segmenterDirectory, "nsfw-seg-breast-x.pt"), + "breast", + ); + await fs.writeFile( + path.join(segmenterDirectory, "nsfw-seg-penis-x.pt"), + "penis", + ); + + await expect(discoverEnhancementModels([root])).resolves.toEqual({ + upscalerModelPath: "", + faceDetectorModelPath: "", + nsfwSegmenterModelPath: "", + }); + }); +}); + +describe("downloadTrustedAsset", () => { + const payload = Buffer.from("trusted model payload"); + const asset: TrustedAsset = { + directory: "upscalers", + filename: "model.pth", + url: "https://models.example/model.pth", + bytes: payload.byteLength, + sha256: createHash("sha256").update(payload).digest("hex"), + }; + + it("promotes a verified download to its final path", async () => { + const root = await temporaryDirectory(); + const target = await downloadTrustedAsset( + root, + asset, + async () => new Response(payload), + ); + + await expect(fs.readFile(target)).resolves.toEqual(payload); + await expect(fs.stat(`${target}.part`)).rejects.toMatchObject({ + code: "ENOENT", + }); + }); + + it("rejects a corrupt download without leaving model files", async () => { + const root = await temporaryDirectory(); + const target = path.join(root, asset.directory, asset.filename); + + await expect( + downloadTrustedAsset( + root, + asset, + async () => new Response("not the expected model"), + ), + ).rejects.toThrow("integrity check"); + await expect(fs.stat(target)).rejects.toMatchObject({ code: "ENOENT" }); + await expect(fs.stat(`${target}.part`)).rejects.toMatchObject({ + code: "ENOENT", + }); + }); +}); diff --git a/electron/enhancement-models.ts b/electron/enhancement-models.ts new file mode 100644 index 0000000..123885c --- /dev/null +++ b/electron/enhancement-models.ts @@ -0,0 +1,199 @@ +import { createHash } from "node:crypto"; +import { promises as fs } from "node:fs"; +import path from "node:path"; +import type { + EnhancementInstallResult, + EnhancementModelKind, + EnhancementModelPaths, +} from "../src/types"; + +export interface TrustedAsset { + directory: string; + filename: string; + url: string; + bytes: number; + sha256: string; +} + +const TRUSTED_ASSETS: Record = { + upscaler: [ + { + directory: "upscalers", + filename: "4x-UltraSharp.pth", + url: "https://huggingface.co/lokCX/4x-Ultrasharp/resolve/1856559b50de25116a7c07261177dd128f1f5664/4x-UltraSharp.pth", + bytes: 66_961_958, + sha256: + "a5812231fc936b42af08a5edba784195495d303d5b3248c24489ef0c4021fe01", + }, + ], + faceDetector: [ + { + directory: "face_detector", + filename: "face_yolov8n.pt", + url: "https://huggingface.co/Bingsu/adetailer/resolve/53cc19de382014514d9d4038601d261a7faa9b7b/face_yolov8n.pt", + bytes: 6_230_011, + sha256: + "70b640f8f60b1cf0dcc72f30caf3da9495eb2fb6509da48c53374ad6806e6a9c", + }, + ], + nsfwSegmenter: [ + { + directory: "nsfw_segmentation", + filename: "nsfw-seg-breast-x.pt", + url: "https://huggingface.co/NSFW-API/NSFW_Segmentation/resolve/818199566a6e3aa0acec2dd1f98bc98c9b99750d/nsfw-seg-breast-x.pt", + bytes: 124_800_929, + sha256: + "5ad882ddaf149873be131943b373da9f15a0603c91508e2291ece83d729c8ecc", + }, + { + directory: "nsfw_segmentation", + filename: "nsfw-seg-penis-x.pt", + url: "https://huggingface.co/NSFW-API/NSFW_Segmentation/resolve/818199566a6e3aa0acec2dd1f98bc98c9b99750d/nsfw-seg-penis-x.pt", + bytes: 124_844_641, + sha256: + "49d9fc8ee67d3bdee44e46bf75aeb76058f5cd074bac027b55ff071b63a32e21", + }, + { + directory: "nsfw_segmentation", + filename: "nsfw-seg-vagina-x.pt", + url: "https://huggingface.co/NSFW-API/NSFW_Segmentation/resolve/818199566a6e3aa0acec2dd1f98bc98c9b99750d/nsfw-seg-vagina-x.pt", + bytes: 124_817_249, + sha256: + "f8349ae348f5cf041809c38bf38d2c80a0f3dccdb1375b97971358cde742197f", + }, + ], +}; + +const UPSCALER_FILES = [ + ["upscalers", "4x-UltraSharp.pth"], + ["4x-UltraSharp.pth"], +]; +const FACE_DETECTOR_FILES = [ + ["face_detector", "face_yolov8n.pt"], + ["face_yolov8n.pt"], +]; +const NSFW_SEGMENTER_FILES = [ + "nsfw-seg-breast-x.pt", + "nsfw-seg-penis-x.pt", + "nsfw-seg-vagina-x.pt", +]; + +async function firstUsableFile( + roots: string[], + candidates: string[][], +): Promise { + for (const root of roots) { + for (const candidate of candidates) { + const filePath = path.resolve(root, ...candidate); + try { + const stat = await fs.stat(filePath); + if (stat.isFile() && stat.size > 0) return filePath; + } catch { + // Continue through known locations. + } + } + } + return ""; +} + +async function firstUsableSegmenterDirectory(roots: string[]): Promise { + for (const root of roots) { + const directory = path.resolve(root, "nsfw_segmentation"); + const filesAreUsable = await Promise.all( + NSFW_SEGMENTER_FILES.map(async (filename) => { + try { + const stat = await fs.stat(path.join(directory, filename)); + return stat.isFile() && stat.size > 0; + } catch { + return false; + } + }), + ); + if (filesAreUsable.every(Boolean)) return directory; + } + return ""; +} + +export async function discoverEnhancementModels( + modelRoots: string[], +): Promise { + const [upscalerModelPath, faceDetectorModelPath, nsfwSegmenterModelPath] = + await Promise.all([ + firstUsableFile(modelRoots, UPSCALER_FILES), + firstUsableFile(modelRoots, FACE_DETECTOR_FILES), + firstUsableSegmenterDirectory(modelRoots), + ]); + return { + upscalerModelPath, + faceDetectorModelPath, + nsfwSegmenterModelPath, + }; +} + +function validPayload(payload: Buffer, asset: TrustedAsset): boolean { + return ( + payload.byteLength === asset.bytes && + createHash("sha256").update(payload).digest("hex") === asset.sha256 + ); +} + +export async function downloadTrustedAsset( + modelsRoot: string, + asset: TrustedAsset, + fetchAsset: (url: string) => Promise, +): Promise { + const targetDirectory = path.resolve(modelsRoot, asset.directory); + const target = path.join(targetDirectory, asset.filename); + const temporary = `${target}.part`; + + try { + const existing = await fs.readFile(target); + if (validPayload(existing, asset)) return target; + } catch (error) { + if ((error as NodeJS.ErrnoException).code !== "ENOENT") throw error; + } + + const response = await fetchAsset(asset.url); + if (!response.ok) { + throw new Error( + `Model download failed with HTTP ${response.status} ${response.statusText}.`, + ); + } + const payload = Buffer.from(await response.arrayBuffer()); + if (!validPayload(payload, asset)) { + throw new Error("Downloaded model failed its integrity check."); + } + + await fs.mkdir(targetDirectory, { recursive: true }); + try { + await fs.writeFile(temporary, payload); + await fs.rm(target, { force: true }); + await fs.rename(temporary, target); + } catch (error) { + await fs.rm(temporary, { force: true }); + throw error; + } + return target; +} + +export async function installEnhancementModel( + modelsRoot: string, + kind: EnhancementModelKind, + fetchAsset: (url: string) => Promise, +): Promise { + const assets = TRUSTED_ASSETS[kind]; + if (!assets) throw new Error("Unknown enhancement model."); + const installedPaths: string[] = []; + for (const asset of assets) { + installedPaths.push( + await downloadTrustedAsset(modelsRoot, asset, fetchAsset), + ); + } + return { + kind, + path: + kind === "nsfwSegmenter" + ? path.dirname(installedPaths[0]) + : installedPaths[0], + }; +} diff --git a/electron/jobs.test.ts b/electron/jobs.test.ts index cdf135d..3978e68 100644 --- a/electron/jobs.test.ts +++ b/electron/jobs.test.ts @@ -51,10 +51,20 @@ describe("LocalJobManager", () => { const modelPath = path.join(root, "image-model"); const upscalerPath = path.join(root, "4x-UltraSharp.pth"); const faceDetectorPath = path.join(root, "face-detector.pt"); + const segmenterPath = path.join(root, "nsfw-segmentation"); await fs.mkdir(modelPath); + await fs.mkdir(segmenterPath); await fs.writeFile(path.join(modelPath, "model_index.json"), "{}"); await fs.writeFile(upscalerPath, "fixture"); await fs.writeFile(faceDetectorPath, "fixture"); + await Promise.all( + ["breast", "penis", "vagina"].map((region) => + fs.writeFile( + path.join(segmenterPath, `nsfw-seg-${region}-x.pt`), + "fixture", + ), + ), + ); const events: JobEvent[] = []; const manager = new LocalJobManager({ runtimeDirectory: () => root, @@ -87,6 +97,8 @@ describe("LocalJobManager", () => { upscale: true, upscaleFactor: 4, upscalerModelPath: upscalerPath, + nsfwSegmentation: true, + nsfwSegmenterModelPath: segmenterPath, }, (event) => events.push(event), ); @@ -107,7 +119,7 @@ describe("LocalJobManager", () => { root, "outputs", "images", - "fixture-preview-1.png", + "preview-fixture-1.png", ), step: 1, total: 2, @@ -134,6 +146,12 @@ describe("LocalJobManager", () => { height: 512, }), ); + await expect( + fs.stat(path.join(root, "outputs", "images", "preview-fixture-1.png")), + ).rejects.toMatchObject({ code: "ENOENT" }); + await expect( + fs.stat(path.join(root, "outputs", "images", "fixture.png")), + ).resolves.toMatchObject({ size: 7 }); await expect( fs.readFile(path.join(root, "outputs", "images", "request.json"), "utf8"), ).resolves.toEqual( @@ -151,10 +169,31 @@ describe("LocalJobManager", () => { upscale: true, upscale_factor: 4, upscaler_model: upscalerPath, + nsfw_segmentation: true, + nsfw_segmenter_model_dir: segmenterPath, }), ); }); + it("purges stale previews without deleting final images", async () => { + const root = await temporaryDirectory(); + const outputDirectory = path.join(root, "outputs", "images"); + const preview = path.join(outputDirectory, "preview-stale-job-001.png"); + const final = path.join(outputDirectory, "final-image.png"); + await fs.mkdir(outputDirectory, { recursive: true }); + await fs.writeFile(preview, "preview"); + await fs.writeFile(final, "final"); + const manager = new LocalJobManager({ + runtimeDirectory: () => root, + outputDirectory: () => path.join(root, "outputs"), + }); + + await manager.cleanupPreviews(); + + await expect(fs.stat(preview)).rejects.toMatchObject({ code: "ENOENT" }); + await expect(fs.readFile(final, "utf8")).resolves.toBe("final"); + }); + it("runs training only with a local Transformers model and dataset", async () => { const root = await temporaryDirectory(); const modelPath = path.join(root, "language-model"); @@ -228,6 +267,8 @@ describe("LocalJobManager", () => { upscale: false, upscaleFactor: 2 as const, upscalerModelPath: "", + nsfwSegmentation: false, + nsfwSegmenterModelPath: "", }; await expect( @@ -250,6 +291,16 @@ describe("LocalJobManager", () => { () => undefined, ), ).rejects.toThrow("Upscaler model was not found"); + await expect( + manager.startImage( + { + ...request, + nsfwSegmentation: true, + nsfwSegmenterModelPath: path.join(root, "missing-segmenters"), + }, + () => undefined, + ), + ).rejects.toThrow("NSFW segmentation model directory was not found"); }); it("rejects single-file image checkpoints before spawning Python", async () => { @@ -285,6 +336,8 @@ describe("LocalJobManager", () => { upscale: false, upscaleFactor: 2, upscalerModelPath: "", + nsfwSegmentation: false, + nsfwSegmenterModelPath: "", }, () => undefined, ), diff --git a/electron/jobs.ts b/electron/jobs.ts index 25d9f8c..4361264 100644 --- a/electron/jobs.ts +++ b/electron/jobs.ts @@ -19,9 +19,16 @@ interface ActiveJob { kind: JobKind; outputDirectory: string; emit: (event: JobEvent) => void; + previewPaths: Set; } const JOB_ID_PATTERN = /^[A-Za-z0-9_-]{1,120}$/; +const PREVIEW_FILE_PATTERN = /^preview-[A-Za-z0-9_-]+\.png$/; +const NSFW_SEGMENTER_FILES = [ + "nsfw-seg-breast-x.pt", + "nsfw-seg-penis-x.pt", + "nsfw-seg-vagina-x.pt", +]; function assertNumber( value: number, @@ -122,6 +129,7 @@ export class LocalJobManager { assertBoolean(request.faceFix, "Face Fix"); assertNumber(request.faceFixStrength, "Face Fix strength", 0.1, 0.8); assertBoolean(request.upscale, "Upscale"); + assertBoolean(request.nsfwSegmentation, "NSFW segmentation"); if (request.upscaleFactor !== 2 && request.upscaleFactor !== 4) { throw new Error("Upscale factor must be 2 or 4."); } @@ -139,6 +147,20 @@ export class LocalJobManager { ".safetensors", ]); } + if (request.nsfwSegmentation) { + await assertDirectory( + request.nsfwSegmenterModelPath, + "NSFW segmentation model directory", + ); + await Promise.all( + NSFW_SEGMENTER_FILES.map((filename) => + assertFile( + path.join(request.nsfwSegmenterModelPath, filename), + `NSFW segmentation model ${filename}`, + ), + ), + ); + } const outputDirectory = path.join(this.paths.outputDirectory(), "images"); await fs.mkdir(outputDirectory, { recursive: true }); @@ -173,6 +195,8 @@ export class LocalJobManager { upscale: request.upscale, upscale_factor: request.upscaleFactor, upscaler_model: request.upscalerModelPath, + nsfw_segmentation: request.nsfwSegmentation, + nsfw_segmenter_model_dir: request.nsfwSegmenterModelPath, }, outputDirectory, emit, @@ -244,6 +268,26 @@ export class LocalJobManager { this.processes.cancelAll(); } + async cleanupPreviews(): Promise { + const outputDirectory = path.join(this.paths.outputDirectory(), "images"); + let entries; + try { + entries = await fs.readdir(outputDirectory, { withFileTypes: true }); + } catch (error) { + if ((error as NodeJS.ErrnoException).code === "ENOENT") return; + throw error; + } + await Promise.all( + entries + .filter( + (entry) => entry.isFile() && PREVIEW_FILE_PATTERN.test(entry.name), + ) + .map((entry) => + fs.rm(path.join(outputDirectory, entry.name), { force: true }), + ), + ); + } + private assertPythonPath(value: string): void { if ( !value.trim() || @@ -266,7 +310,12 @@ export class LocalJobManager { ): void { if (this.active.has(jobId)) throw new Error(`Job ${jobId} is already active.`); - this.active.set(jobId, { kind, outputDirectory, emit }); + this.active.set(jobId, { + kind, + outputDirectory, + emit, + previewPaths: new Set(), + }); emit({ jobId, kind, type: "queued", progress: 0 }); try { this.processes.start( @@ -315,14 +364,23 @@ export class LocalJobManager { return; } if (event.type === "preview") { + let outputPath = ""; try { if (job.kind !== "image" || typeof event.path !== "string") { throw new Error("Worker preview is not a valid image output."); } - const outputPath = path.resolve(event.path); + outputPath = path.resolve(event.path); assertInside(job.outputDirectory, outputPath); + if (!PREVIEW_FILE_PATTERN.test(path.basename(outputPath))) { + throw new Error("Worker preview has an invalid filename."); + } + job.previewPaths.add(outputPath); const stat = await fs.stat(outputPath); if (!stat.isFile()) throw new Error("Worker preview is not a file."); + if (this.active.get(jobId) !== job) { + await fs.rm(outputPath, { force: true }); + return; + } const step = Number(event.step); const total = Number(event.total); const progress = @@ -345,6 +403,8 @@ export class LocalJobManager { : undefined, }); } catch (error) { + if (outputPath) job.previewPaths.delete(outputPath); + if (this.active.get(jobId) !== job) return; job.emit({ jobId, kind: job.kind, @@ -367,6 +427,7 @@ export class LocalJobManager { } if (event.type === "error" || event.type === "cancelled") { this.active.delete(jobId); + await this.cleanupJobPreviews(job); job.emit({ jobId, kind: job.kind, @@ -391,6 +452,7 @@ export class LocalJobManager { throw new Error("Training worker output is not a directory."); } this.active.delete(jobId); + await this.cleanupJobPreviews(job); job.emit({ jobId, kind: job.kind, @@ -406,6 +468,7 @@ export class LocalJobManager { }); } catch (error) { this.active.delete(jobId); + await this.cleanupJobPreviews(job); job.emit({ jobId, kind: job.kind, @@ -415,4 +478,14 @@ export class LocalJobManager { }); } } + + private async cleanupJobPreviews(job: ActiveJob): Promise { + const previewPaths = [...job.previewPaths]; + job.previewPaths.clear(); + await Promise.all( + previewPaths.map((previewPath) => + fs.rm(previewPath, { force: true }).catch(() => undefined), + ), + ); + } } diff --git a/electron/main.ts b/electron/main.ts index 8daf44d..91bc008 100644 --- a/electron/main.ts +++ b/electron/main.ts @@ -20,6 +20,9 @@ import { import type { ChatRequest, ChatStreamEvent, + EnhancementInstallResult, + EnhancementModelKind, + EnhancementModelPaths, ImageAttachment, ImageGenerationRequest, ImageModel, @@ -36,6 +39,10 @@ import type { TrainingRequest, } from "../src/types"; import { runMcpChat } from "./chat"; +import { + discoverEnhancementModels, + installEnhancementModel, +} from "./enhancement-models"; import { LocalJobManager } from "./jobs"; import { closeMcpConnections, disconnectMcpServer, testMcpServer } from "./mcp"; import { discoverModelCatalog } from "./model-catalog"; @@ -175,6 +182,17 @@ function emitJob(sender: WebContents, event: JobEvent): void { }); } +function enhancementModelRoots(): string[] { + const appRoot = app.getAppPath(); + return [ + path.join(app.getPath("userData"), "models"), + path.join(appRoot, "models"), + path.join(process.resourcesPath, "models"), + path.resolve(appRoot, "..", "Lavely-LLM", "models"), + path.resolve(process.cwd(), "..", "Lavely-LLM", "models"), + ].filter((root, index, roots) => roots.indexOf(root) === index); +} + async function imageForOllama(value: string): Promise { if (value.startsWith(`${ATTACHMENT_SCHEME}:`)) { return (await fs.readFile(attachmentFilePath(value))).toString("base64"); @@ -662,6 +680,21 @@ function registerIpc(): void { shell.openExternal(modelCatalogUrl(url)), ); + ipcMain.handle( + "enhancements:discover", + (): Promise => + discoverEnhancementModels(enhancementModelRoots()), + ); + ipcMain.handle( + "enhancements:install", + (_event, kind: EnhancementModelKind): Promise => + installEnhancementModel( + path.join(app.getPath("userData"), "models"), + kind, + (url) => net.fetch(url), + ), + ); + ipcMain.handle("mcp:test-server", (_event, server: McpServerConfig) => testMcpServer(server), ); @@ -846,6 +879,18 @@ function registerIpc(): void { return result.canceled ? null : (result.filePaths[0] ?? null); }, ); + + ipcMain.handle( + "dialog:choose-nsfw-segmenter-models", + async (): Promise => { + const result = await dialog.showOpenDialog({ + title: "Select an NSFW segmentation model directory", + buttonLabel: "Select models", + properties: ["openDirectory"], + }); + return result.canceled ? null : (result.filePaths[0] ?? null); + }, + ); } function createWindow(): void { @@ -891,7 +936,7 @@ function createWindow(): void { registerIpc(); -void app.whenReady().then(() => { +void app.whenReady().then(async () => { app.setAppUserModelId("io.localforge.desktop"); void protocol.handle(ATTACHMENT_SCHEME, async (request) => { try { @@ -911,6 +956,9 @@ void app.whenReady().then(() => { return new Response("Output not found.", { status: 404 }); } }); + await jobs.cleanupPreviews().catch((error) => { + console.warn("Could not remove temporary image previews:", error); + }); createWindow(); app.on("activate", () => { if (BrowserWindow.getAllWindows().length === 0) createWindow(); diff --git a/electron/preload.ts b/electron/preload.ts index ad501f0..29c4005 100644 --- a/electron/preload.ts +++ b/electron/preload.ts @@ -3,6 +3,9 @@ import type { AppInfo, ChatRequest, ChatStreamEvent, + EnhancementInstallResult, + EnhancementModelKind, + EnhancementModelPaths, ForgeApi, ImageAttachment, ImageGenerationRequest, @@ -73,6 +76,17 @@ const api: ForgeApi = { ) as Promise, open: (url: string) => ipcRenderer.invoke("catalog:open", url), }, + enhancements: { + discover: () => + ipcRenderer.invoke( + "enhancements:discover", + ) as Promise, + install: (kind: EnhancementModelKind) => + ipcRenderer.invoke( + "enhancements:install", + kind, + ) as Promise, + }, mcp: { testServer: (server: McpServerConfig) => ipcRenderer.invoke("mcp:test-server", server) as Promise, @@ -123,6 +137,10 @@ const api: ForgeApi = { ipcRenderer.invoke("dialog:choose-face-detector-model") as Promise< string | null >, + chooseNsfwSegmenterModels: () => + ipcRenderer.invoke("dialog:choose-nsfw-segmenter-models") as Promise< + string | null + >, }, }; diff --git a/index.html b/index.html index 25b445c..5dd0833 100644 --- a/index.html +++ b/index.html @@ -4,7 +4,7 @@ diff --git a/runtime/image_worker.py b/runtime/image_worker.py index faf0ecb..9ba47fd 100644 --- a/runtime/image_worker.py +++ b/runtime/image_worker.py @@ -12,6 +12,12 @@ sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding="utf-8", errors="replace") sys.stderr = io.TextIOWrapper(sys.stderr.buffer, encoding="utf-8", errors="replace") +NSFW_SEGMENTER_MODELS = { + "breast": "nsfw-seg-breast-x.pt", + "penis": "nsfw-seg-penis-x.pt", + "vagina": "nsfw-seg-vagina-x.pt", +} + def emit(payload: dict) -> None: print(json.dumps(payload), flush=True) @@ -165,6 +171,54 @@ def require_model_file( return model_path +def require_segmenter_directory(value: object) -> Path: + model_dir = Path(str(value or "")).expanduser().resolve() + if not model_dir.is_dir(): + raise ValueError(f"NSFW segmentation model directory was not found: {model_dir}") + for filename in NSFW_SEGMENTER_MODELS.values(): + if not (model_dir / filename).is_file(): + raise ValueError(f"NSFW segmentation model was not found: {filename}") + return model_dir + + +def segment_nsfw( + model_dir: Path, + image: Any, + conf: float = 0.3, + model_factory: Any | None = None, +) -> dict[str, Any]: + import numpy as np + from PIL import Image + + if model_factory is None: + try: + from ultralytics import YOLO + except ImportError as error: + raise RuntimeError( + "NSFW segmentation requires ultralytics. " + "Reinstall runtime/requirements-image.txt." + ) from error + model_factory = YOLO + + width, height = image.size + masks: dict[str, Any] = {} + for region, filename in NSFW_SEGMENTER_MODELS.items(): + model = model_factory(str(model_dir / filename)) + results = model.predict(image, conf=conf, verbose=False) + combined = np.zeros((height, width), dtype=np.uint8) + for result in results: + if result.masks is None: + continue + for mask_data in result.masks.data: + mask_array = mask_data.cpu().numpy() + mask = Image.fromarray((mask_array * 255).astype(np.uint8)) + mask = mask.resize((width, height), Image.Resampling.NEAREST) + combined = np.maximum(combined, np.asarray(mask)) + if combined.max() > 0: + masks[region] = Image.fromarray(combined) + return masks + + def expanded_face_box( box: tuple[int, int, int, int], image_width: int, @@ -403,6 +457,7 @@ def main() -> None: 0.8, ) upscale = bool(request.get("upscale", False)) + nsfw_segmentation = bool(request.get("nsfw_segmentation", False)) upscale_factor = bounded_int( request.get("upscale_factor", 2), "upscale_factor", 2, 4 ) @@ -540,11 +595,35 @@ def legacy_progress_callback(step_index, _timestep, latents): metadata.add_text("face_fix_strength", str(face_fix_strength)) metadata.add_text("upscale", str(upscale).lower()) metadata.add_text("upscale_factor", str(upscale_factor)) + metadata.add_text("nsfw_segmentation", str(nsfw_segmentation).lower()) metadata.add_text("generator", "Local Forge") image.save(output_path, pnginfo=metadata) except ImportError: image.save(output_path) + if nsfw_segmentation: + try: + segmenter_dir = require_segmenter_directory( + request.get("nsfw_segmenter_model_dir") + ) + emit({"type": "status", "message": "Running NSFW segmentation..."}) + masks = segment_nsfw(segmenter_dir, image) + for region, mask in masks.items(): + mask_path = output_dir / f"{output_path.stem}_mask_{region}.png" + mask.save(mask_path) + if masks: + regions = ", ".join(sorted(masks)) + emit({"type": "log", "message": f"Saved NSFW masks: {regions}."}) + else: + emit( + { + "type": "log", + "message": "NSFW segmentation found no matching regions.", + } + ) + except Exception as error: + emit({"type": "log", "message": f"NSFW segmentation skipped: {error}"}) + emit( { "type": "done", diff --git a/src/App.css b/src/App.css index 0ee01b4..1ea57ff 100644 --- a/src/App.css +++ b/src/App.css @@ -2159,6 +2159,34 @@ font-size: 9px; } +.studio-enhancement-error { + margin: 7px 0 2px; + color: var(--danger); + font-size: 9px; + line-height: 1.4; +} + +.studio-model-install { + display: inline-flex; + min-width: 72px; + min-height: 27px; + align-items: center; + justify-content: center; + gap: 5px; + padding: 5px 8px; + border: 1px solid color-mix(in srgb, var(--ember) 42%, var(--line)); + border-radius: 6px; + background: color-mix(in srgb, var(--ember) 8%, var(--panel-strong)); + color: var(--ember); + font-size: 9px; + font-weight: 700; +} + +.studio-model-install:disabled { + cursor: wait; + opacity: 0.55; +} + .studio-enhancements .toggle-row:first-of-type { margin-top: 4px; } diff --git a/src/App.tsx b/src/App.tsx index 28d2917..c500e03 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -38,6 +38,7 @@ import type { } from "./state/workspace"; import type { AppView, + EnhancementModelKind, ImageModel, JobEvent, OllamaModel, @@ -96,6 +97,75 @@ function App() { const [sidebarCollapsed, setSidebarCollapsed] = useState(false); const [commandOpen, setCommandOpen] = useState(false); const [imagePreview, setImagePreview] = useState(null); + const [installingEnhancement, setInstallingEnhancement] = + useState(null); + const [enhancementInstallError, setEnhancementInstallError] = useState(""); + + const discoverEnhancements = useEffectEvent(async () => { + try { + const discovered = await forgeApi.enhancements.discover(); + setWorkspace((current) => { + const upscalerModelPath = current.settings.upscalerModelPath.trim() + ? current.settings.upscalerModelPath + : discovered.upscalerModelPath; + const faceDetectorModelPath = + current.settings.faceDetectorModelPath.trim() + ? current.settings.faceDetectorModelPath + : discovered.faceDetectorModelPath; + const nsfwSegmenterModelPath = + current.settings.nsfwSegmenterModelPath.trim() + ? current.settings.nsfwSegmenterModelPath + : discovered.nsfwSegmenterModelPath; + if ( + upscalerModelPath === current.settings.upscalerModelPath && + faceDetectorModelPath === current.settings.faceDetectorModelPath && + nsfwSegmenterModelPath === current.settings.nsfwSegmenterModelPath + ) { + return current; + } + return { + ...current, + settings: { + ...current.settings, + upscalerModelPath, + faceDetectorModelPath, + nsfwSegmenterModelPath, + }, + }; + }); + } catch (error) { + console.warn("Could not discover enhancement models:", error); + } + }); + + useEffect(() => { + if (loaded) void discoverEnhancements(); + }, [loaded]); + + async function installEnhancement(kind: EnhancementModelKind) { + setInstallingEnhancement(kind); + setEnhancementInstallError(""); + try { + const installed = await forgeApi.enhancements.install(kind); + setWorkspace((current) => ({ + ...current, + settings: { + ...current.settings, + ...(installed.kind === "upscaler" + ? { upscalerModelPath: installed.path } + : installed.kind === "faceDetector" + ? { faceDetectorModelPath: installed.path } + : { nsfwSegmenterModelPath: installed.path }), + }, + })); + } catch (error) { + setEnhancementInstallError( + error instanceof Error ? error.message : "Model installation failed.", + ); + } finally { + setInstallingEnhancement(null); + } + } async function refreshModels() { try { @@ -350,6 +420,8 @@ function App() { upscale: Boolean(recipe.upscale), upscaleFactor: recipe.upscaleFactor ?? 2, upscalerModelPath: workspace.settings.upscalerModelPath, + nsfwSegmentation: Boolean(recipe.nsfwSegmentation), + nsfwSegmenterModelPath: workspace.settings.nsfwSegmenterModelPath, }); } catch (error) { failRun(jobId, error); @@ -424,6 +496,9 @@ function App() { workspace.studio.upscale && Boolean(workspace.settings.upscalerModelPath.trim()), upscaleFactor: workspace.studio.upscaleFactor, + nsfwSegmentation: + workspace.studio.nsfwSegmentation && + Boolean(workspace.settings.nsfwSegmenterModelPath.trim()), }; await queueImageRun( recipe, @@ -500,6 +575,7 @@ function App() { faceFixStrength: recipe.faceFixStrength ?? 0.45, upscale: recipe.upscale ?? false, upscaleFactor: recipe.upscaleFactor ?? 2, + nsfwSegmentation: recipe.nsfwSegmentation ?? false, activeAsset: outputUrl, }, })); @@ -759,6 +835,11 @@ function App() { faceDetectorConfigured={Boolean( workspace.settings.faceDetectorModelPath.trim(), )} + nsfwSegmenterConfigured={Boolean( + workspace.settings.nsfwSegmenterModelPath.trim(), + )} + installingEnhancement={installingEnhancement} + enhancementInstallError={enhancementInstallError} preview={ imagePreview && imagePreview.jobId === activeImageRun?.id && @@ -782,6 +863,7 @@ function App() { onImportAssets={importStudioAssets} onRemoveModel={removeImageModel} onOpenSettings={() => selectView("settings")} + onInstallEnhancement={(kind) => void installEnhancement(kind)} /> )} {workspace.activeView === "library" && ( diff --git a/src/features/operations/OperationalViews.tsx b/src/features/operations/OperationalViews.tsx index c480c93..af238bb 100644 --- a/src/features/operations/OperationalViews.tsx +++ b/src/features/operations/OperationalViews.tsx @@ -355,6 +355,9 @@ function RunRecipe({ run }: { run: ForgeRun }) { ...(recipe.upscale ? [["Upscale", `${recipe.upscaleFactor ?? 2}x`]] : []), + ...(recipe.nsfwSegmentation + ? [["NSFW segmentation", "Save detected-region masks"]] + : []), ["Prompt", recipe.prompt], ...(recipe.negativePrompt ? [["Negative prompt", recipe.negativePrompt]] @@ -1139,6 +1142,11 @@ export function SettingsView({ if (selected) onChange({ faceDetectorModelPath: selected }); } + async function browseNsfwSegmenterModels() { + const selected = await forgeApi.dialog.chooseNsfwSegmenterModels(); + if (selected) onChange({ nsfwSegmenterModelPath: selected }); + } + return (
@@ -1410,9 +1418,33 @@ export function SettingsView({ +

- Model files stay local. Local Forge does not download or - bundle third-party enhancement weights. + Install actions download pinned, checksum-verified weights + to Local Forge storage. You can also choose compatible local + files manually.

diff --git a/src/features/studio/StudioView.tsx b/src/features/studio/StudioView.tsx index f7bb44d..c992f85 100644 --- a/src/features/studio/StudioView.tsx +++ b/src/features/studio/StudioView.tsx @@ -1,4 +1,5 @@ import { + Download, FolderSearch, Image as ImageIcon, ImagePlus, @@ -18,7 +19,7 @@ import { type LibraryAsset, type StudioState, } from "../../state/workspace"; -import type { ImageModel } from "../../types"; +import type { EnhancementModelKind, ImageModel } from "../../types"; const stylePresets = [ { name: "Editorial", description: "Natural light / tactile" }, @@ -34,6 +35,9 @@ interface StudioViewProps { activeRun?: ForgeRun; upscalerConfigured: boolean; faceDetectorConfigured: boolean; + nsfwSegmenterConfigured: boolean; + installingEnhancement: EnhancementModelKind | null; + enhancementInstallError: string; preview?: { src: string; step?: number; @@ -46,6 +50,7 @@ interface StudioViewProps { onImportAssets: () => Promise; onRemoveModel: (model: ImageModel) => void; onOpenSettings: () => void; + onInstallEnhancement: (kind: EnhancementModelKind) => void; } export function StudioView({ @@ -55,6 +60,9 @@ export function StudioView({ activeRun, upscalerConfigured, faceDetectorConfigured, + nsfwSegmenterConfigured, + installingEnhancement, + enhancementInstallError, preview, onChange, onGenerate, @@ -63,6 +71,7 @@ export function StudioView({ onImportAssets, onRemoveModel, onOpenSettings, + onInstallEnhancement, }: StudioViewProps) { const [zoom, setZoom] = useState(67); const [notice, setNotice] = useState(""); @@ -440,13 +449,20 @@ export function StudioView({
Enhance - {(!faceDetectorConfigured || !upscalerConfigured) && ( + {(!faceDetectorConfigured || + !upscalerConfigured || + !nsfwSegmenterConfigured) && ( )}
-