Ansel 0.0
A darktable fork - bloat + design vision
Loading...
Searching...
No Matches
nn_model_test.c
Go to the documentation of this file.
1/*
2 This file is part of Ansel,
3 Copyright (C) 2026 Aurélien PIERRE.
4
5 Ansel is free software: you can redistribute it and/or modify
6 it under the terms of the GNU General Public License as published by
7 the Free Software Foundation, either version 3 of the License, or
8 (at your option) any later version.
9
10 Ansel is distributed in the hope that it will be useful,
11 but WITHOUT ANY WARRANTY; without even the implied warranty of
12 MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
13 GNU General Public License for more details.
14
15 You should have received a copy of the GNU General Public License
16 along with Ansel. If not, see <http://www.gnu.org/licenses/>.
17*/
18
19/* Golden-fixture parity test for the .anselnn loader + CPU U-Net executor.
20 * The fixture is produced by scripts/make_fixture.py in the ansel-denoise
21 * training repo from the same model file; the C output must match the torch
22 * reference within the stated absolute tolerance.
23 *
24 * Standalone build (no ansel build system needed):
25 * gcc -O2 -fopenmp -Isrc src/common/nn_model.c tests/nn_model_test.c \
26 * $(pkg-config --cflags --libs json-glib-1.0) -lm -o nn_model_test
27 * Usage: nn_model_test <model.anselnn> <fixture-dir> [N] [--xtrans]
28 * N defaults to 96. --xtrans selects the 6-px superpixel bin for a fixture
29 * generated with `make_fixture.py --cfa xtrans`; without it the 4-px Bayer
30 * bin is assumed, and a mismatch shows up immediately as a binning-contract
31 * failure rather than as a silently wrong comparison.
32 */
33
34#include "common/nn_model.h"
35
36#include <math.h>
37#include <stdio.h>
38#include <stdlib.h>
39#include <string.h>
40#include <json-glib/json-glib.h>
41#include <time.h>
42
43static double max_abs_diff(const float *a, const float *b, size_t count)
44{
45 double max_abs = 0.0;
46 for(size_t i = 0; i < count; i++)
47 {
48 const double d = fabs((double)a[i] - b[i]);
49 if(d > max_abs) max_abs = d;
50 }
51 return max_abs;
52}
53
54static float *read_f32(const char *dir, const char *name, size_t count)
55{
56 char path[1024];
57 snprintf(path, sizeof(path), "%s/%s", dir, name);
58 FILE *f = fopen(path, "rb");
59 if(!f)
60 {
61 fprintf(stderr, "cannot open %s\n", path);
62 return NULL;
63 }
64 float *buf = malloc(count * sizeof(float));
65 const size_t got = buf ? fread(buf, sizeof(float), count, f) : 0;
66 fclose(f);
67 if(got != count)
68 {
69 fprintf(stderr, "%s: expected %zu floats, got %zu\n", path, count, got);
70 free(buf);
71 return NULL;
72 }
73 return buf;
74}
75
76/* Read what the fixture says about itself: which CFA it was generated for (so the superpixel
77 * bin factor is not a flag the caller has to remember) and which model it was generated FROM.
78 *
79 * The hash check is the important half. A fixture pins model_sha256, but nothing used to
80 * enforce it, so running against a model the fixture was not built from produced a large error
81 * that reads exactly like a broken executor -- which is how six perfectly good models once
82 * looked like six failures. Refuse to run instead. */
83static int read_meta(const char *dir, const char *model_path, int *is_xtrans)
84{
85 char path[4096];
86 snprintf(path, sizeof(path), "%s/fixture-meta.json", dir);
87
89 GError *error = NULL;
91 {
92 fprintf(stderr, "cannot read %s: %s\n", path, error ? error->message : "?");
95 return 1;
96 }
98
99 if(json_object_has_member(root, "cfa"))
100 *is_xtrans = !g_strcmp0(json_object_get_string_member(root, "cfa"), "xtrans");
101
102 int rc = 0;
103 if(json_object_has_member(root, "model_sha256"))
104 {
105 const char *want = json_object_get_string_member(root, "model_sha256");
106 gchar *blob = NULL;
107 gsize len = 0;
108 if(g_file_get_contents(model_path, &blob, &len, NULL))
109 {
110 gchar *got = g_compute_checksum_for_data(G_CHECKSUM_SHA256, (const guchar *)blob, len);
111 if(g_strcmp0(got, want))
112 {
114 "FAIL: this fixture was generated from a different model.\n"
115 " fixture pins %s\n model is %s\n"
116 " Regenerate with scripts/make_fixture.py in the ansel-denoise repo.\n",
117 want, got);
118 rc = 1;
119 }
120 g_free(got);
121 g_free(blob);
122 }
123 }
125 return rc;
126}
127
128int main(int argc, char *argv[])
129{
130 if(argc < 3)
131 {
132 fprintf(stderr, "usage: %s <model.anselnn> <fixture-dir> [N] [--xtrans]\n", argv[0]);
133 return 2;
134 }
135 const int n = (argc > 3 && argv[3][0] != '-') ? atoi(argv[3]) : 96;
136 /* the fixture declares its own CFA; the flags stay as an override */
137 int is_xtrans = 0;
138 if(read_meta(argv[2], argv[1], &is_xtrans)) return 1;
139 for(int i = 3; i < argc; i++)
140 {
141 if(!strcmp(argv[i], "--xtrans")) is_xtrans = 1;
142 if(!strcmp(argv[i], "--bayer")) is_xtrans = 0;
143 }
144
145 char err[256] = "";
146 dt_nn_model_t *model = dt_nn_model_load(argv[1], err, sizeof(err));
147 if(!model)
148 {
149 fprintf(stderr, "model load failed: %s\n", err);
150 return 2;
151 }
152 printf("model loaded: in=%d out=%d alignment=%d, scratch for %dx%d: %.1f MB\n", dt_nn_model_in_channels(model),
154 dt_nn_unet_scratch_bytes(model, n, n) / 1048576.0);
155
156 const size_t plane = (size_t)n * n;
157 float *in = read_f32(argv[2], "fixture-input.f32", plane * dt_nn_model_in_channels(model));
158 float *expected = read_f32(argv[2], "fixture-expected.f32", plane);
159 float *out = calloc(plane, sizeof(float));
160 if(!in || !expected || !out) return 2;
161
162 /* multi-scale model: gate the binning contract and the coarse stage before
163 * the fine parity below (which runs on the fixture's torch-built guide). */
164 const int bin = dt_nn_model_bin(model, is_xtrans);
165 printf("fixture CFA: %s, superpixel bin %d\n", is_xtrans ? "xtrans" : "bayer", bin);
166 if(bin > 1)
167 {
168 const int cn = n / bin;
169 const size_t cplane = (size_t)cn * cn;
172 float *base = read_f32(argv[2], "fixture-base-planes.f32", plane * 5);
173 float *c_in_exp = read_f32(argv[2], "fixture-coarse-input.f32", cplane * c_in);
174 float *c_out_exp = read_f32(argv[2], "fixture-coarse-expected.f32", cplane * c_out);
175 float *rgb = malloc(cplane * 3 * sizeof(float));
176 float *cnt = malloc(cplane * 3 * sizeof(float));
177 float *c_out_got = malloc(cplane * c_out * sizeof(float));
178 if(!base || !c_in_exp || !c_out_exp || !rgb || !cnt || !c_out_got) return 2;
179
180 /* 1. binning contract: our RGB means vs torch's binned planes 0..2 */
181 dt_nn_bin_planes(base, n, n, bin, rgb, cnt);
182 const double bin_err = max_abs_diff(rgb, c_in_exp, cplane * 3);
183 printf("binning contract: max abs err %.3g (tolerance 1e-6)\n", bin_err);
184 if(bin_err > 1e-6)
185 {
186 fprintf(stderr, "FAIL: binning contract\n");
187 return 1;
188 }
189
190 /* 2. coarse stage parity on torch's own input */
192 {
193 fprintf(stderr, "coarse stage apply failed\n");
194 return 2;
195 }
196 const double c_err = max_abs_diff(c_out_got, c_out_exp, cplane * c_out);
197 printf("coarse stage parity: max abs err %.3g (tolerance 2e-4)\n", c_err);
198 if(c_err > 2e-4)
199 {
200 fprintf(stderr, "FAIL: coarse stage parity\n");
201 return 1;
202 }
203
204 /* 3. end-to-end: our binning -> our coarse -> our guide injection -> fine,
205 * compared against the torch final output (looser: coarse error propagates) */
207 float *fine_in = malloc(plane * fine_in_ch * sizeof(float));
208 float *e2e_out = malloc(plane * sizeof(float));
209 if(!fine_in || !e2e_out) return 2;
210 memcpy(fine_in, base, plane * 5 * sizeof(float));
211 /* rebuild the coarse input from our own binning + the fixture's sigma
212 * planes (positions 3..5 of the coarse input are the binned sigma, which
213 * needs the profile constants — reuse torch's, the contract test above
214 * already pinned our RGB planes) */
218 {
219 fprintf(stderr, "fine stage apply failed\n");
220 return 2;
221 }
222 const double e2e_err = max_abs_diff(e2e_out, expected, plane);
223 printf("end-to-end parity: max abs err %.3g (tolerance 5e-4)\n", e2e_err);
224 if(e2e_err > 5e-4)
225 {
226 fprintf(stderr, "FAIL: end-to-end parity\n");
227 return 1;
228 }
229 free(base);
230 free(c_in_exp);
231 free(c_out_exp);
232 free(rgb);
233 free(cnt);
234 free(c_out_got);
235 free(fine_in);
236 free(e2e_out);
237 }
238
239 struct timespec t0, t1;
241 const int rc = dt_nn_unet_apply(model, in, out, n, n);
243 if(rc)
244 {
245 fprintf(stderr, "dt_nn_unet_apply failed (%d)\n", rc);
246 return 2;
247 }
248 const double ms = (t1.tv_sec - t0.tv_sec) * 1e3 + (t1.tv_nsec - t0.tv_nsec) / 1e6;
249
250 double max_abs = 0.0, sum_sq = 0.0;
251 size_t worst = 0;
252 for(size_t i = 0; i < plane; i++)
253 {
254 const double d = fabs((double)out[i] - expected[i]);
255 if(d > max_abs)
256 {
257 max_abs = d;
258 worst = i;
259 }
260 sum_sq += d * d;
261 }
262 const double rms = sqrt(sum_sq / plane);
263 const double tolerance = 2e-4;
264 printf("parity vs torch: max abs err %.3g (at %zu: %.6f vs %.6f), rms %.3g | %.1f ms for %dx%d\n", max_abs,
265 worst, out[worst], expected[worst], rms, ms, n, n);
266
268 free(in);
269 free(expected);
270 free(out);
271 if(max_abs > tolerance)
272 {
273 fprintf(stderr, "FAIL: max abs err %.3g > %.3g\n", max_abs, tolerance);
274 return 1;
275 }
276 printf("PASS\n");
277 return 0;
278}
static void error(char *msg)
Definition ashift_lsd.c:202
const float f
static dt_aligned_pixel_t rgb
const dt_colormatrix_t dt_aligned_pixel_t out
const char * model
void dt_nn_model_free(dt_nn_model_t *m)
Definition nn_model.c:405
int dt_nn_model_in_channels(const dt_nn_model_t *m)
Definition nn_model.c:416
int dt_nn_model_coarse_in_channels(const dt_nn_model_t *m)
Definition nn_model.c:432
size_t dt_nn_unet_scratch_bytes(const dt_nn_model_t *m, int width, int height)
Definition nn_model.c:736
int dt_nn_unet_apply_stage(const dt_nn_model_t *m, int stage, const float *in, float *out, int width, int height, int apply_residual)
Definition nn_model.c:1009
dt_nn_model_t * dt_nn_model_load(const char *path, char *err, size_t err_len)
Definition nn_model.c:237
__DT_CLONE_TARGETS__ void dt_nn_upsample_nearest(const float *in, int ch, int w, int h, int factor, float *out)
Definition nn_model.c:1058
int dt_nn_model_coarse_out_channels(const dt_nn_model_t *m)
Definition nn_model.c:437
int dt_nn_model_alignment(const dt_nn_model_t *m)
Definition nn_model.c:460
int dt_nn_unet_apply(const dt_nn_model_t *m, const float *in, float *out, int width, int height)
Definition nn_model.c:1004
__DT_CLONE_TARGETS__ void dt_nn_bin_planes(const float *planes, int pw, int ph, int bin, float *out_rgb, float *out_cnt)
Definition nn_model.c:1022
int dt_nn_model_bin(const dt_nn_model_t *m, const int is_xtrans)
Definition nn_model.c:426
int dt_nn_model_out_channels(const dt_nn_model_t *m)
Definition nn_model.c:421
static float * read_f32(const char *dir, const char *name, size_t count)
static double max_abs_diff(const float *a, const float *b, size_t count)
static int read_meta(const char *dir, const char *model_path, int *is_xtrans)
const char * name
Definition pdf.h:90
int main()
Definition prova.c:47