use arithmetic mean for PPISP exposure centering

This commit is contained in:
Harry Chen
2026-09-14 22:04:56 -04:00
parent 6b640051e0
commit 2080089a7e
14 changed files with 8121 additions and 7362 deletions
+19 -11
View File
@@ -718,34 +718,39 @@ static void check_cuda_runtime() {
}
#endif // SS_BACKEND_VULKAN
// PPISP exposure seeds: mean-relative EXIF EV x 0.5 per POST-split slot; empty
// when no image has the tags. The 0.5: PPISP multiplies the sRGB-encoded
// PPISP exposure seeds: 0.5 x EXIF EV per POST-split slot, centred like the
// exposure-mean regularizer; empty without tags. 0.5: PPISP scales the sRGB
// render, where a bracketed +1 EV measures x2^0.49 (0.34-0.76 by tone curve).
static std::vector<float> exif_exposure_evs(const ParsedDataset& ds,
const PostSplitCameras& post,
bool arithmetic_mean,
int& n_found) {
int64_t n = ds.num_cameras;
std::vector<double> ev(n, 0.0);
std::vector<double> gain(n, 0.0);
std::vector<char> has(n, 0);
double sum = 0.0;
n_found = 0;
for (int64_t i = 0; i < n; i++) {
double v;
if (sfm::exifExposureEv(sfm::readExif(ds.image_filenames[i]), v)) {
ev[i] = v;
gain[i] = 0.5 * v;
has[i] = 1;
sum += v;
// sum += std::exp2(v);
sum += gain[i];
n_found++;
}
}
if (n_found == 0) return {};
double mean = sum / n_found;
// double mean = std::log2(sum / n_found);
double center = sum / n_found;
if (arithmetic_mean) {
double sum_exp = 0.0;
for (int64_t i = 0; i < n; i++)
if (has[i]) sum_exp += std::exp2(gain[i] - center);
center += std::log2(sum_exp / n_found);
}
std::vector<float> out((size_t)post.n_post, 0.0f);
for (int64_t i = 0; i < n; i++) {
if (!has[i]) continue;
float v = 0.5f * (float)(ev[i] - mean);
float v = (float)(gain[i] - center);
if (post.K_per_camera.empty()) {
out[i] = v;
} else {
@@ -939,13 +944,16 @@ void TrainerSession::setup_engine() {
std::vector<float> exif_ev;
if (cfg.ppisp_exposure_from_exif) {
int n_exif = 0;
exif_ev = exif_exposure_evs(ds, post, n_exif);
exif_ev = exif_exposure_evs(ds, post,
cfg.ppisp_exposure_arithmetic_mean,
n_exif);
if (n_exif > 0)
log(lfmt(lmsg::ppisp_exif_exposure,
{(long long)n_exif, (long long)ds.num_cameras}));
}
engine_init_ppisp(n_grids, cfg.ppisp_param_type,
cfg.use_adagrad_ppisp_optim, exif_ev);
cfg.use_adagrad_ppisp_optim,
cfg.ppisp_exposure_arithmetic_mean, exif_ev);
st.ppisp_init = true;
}
@@ -201,7 +201,7 @@ static int64_t run_case(const char* which, const char* label,
ttv(s.intr.data(), 4, {C, 4}),
ttv(s.dist.data(), 4, {C, kCameraDistortionParams}));
engine_init_bilagrid_rgb(C, "ppisp", 8, 16, 16, 8, 16, true);
engine_init_ppisp(C, "no_crf", true);
engine_init_ppisp(C, "no_crf", true, true);
}
step_n(s, 1, 6);
Health before = count_dead();
@@ -229,7 +229,8 @@ static bool ppisp_grad_case(float poison) {
ttv(s.vm.data(), 4, {C, 4, 4}),
ttv(s.intr.data(), 4, {C, 4}),
ttv(s.dist.data(), 4, {C, kCameraDistortionParams}));
engine_init_ppisp(C, "no_crf", /*use_adagrad=*/true);
engine_init_ppisp(C, "no_crf", /*use_adagrad=*/true,
/*exposure_arithmetic_mean=*/true);
step_n(s, 1, 2);
int64_t n = engine().ppisp.params.numel();
+59 -10
View File
@@ -88,7 +88,7 @@ int main(int argc, char** argv) {
}
const bool dumping = std::strcmp(argv[1], "dump") == 0;
Rng r(260719u);
Rng r(260719u), r_arith(260914u);
const TypeInfo types[6] = {
{"original", 36, (int)RawPPISPRegLossIndex::length},
@@ -148,34 +148,83 @@ int main(int argc, char** argv) {
// Regularization forward: per-image raw rows are deterministic
// (tight); the summed tail row and the weighted losses derived from
// it are atomic-order dependent (loose).
{
for (int arith = 0; arith < 2; arith++) {
// Own stream for the added mode, so every later block keeps its data.
Rng& rr = arith ? r_arith : r;
const int64_t Bp = 6;
float* params = upload(r.vec(Bp * t.n_params, -0.5f, 0.5f));
std::vector<float> params_h = rr.vec(Bp * t.n_params, -0.5f, 0.5f);
float* params = upload(params_h);
std::array<float, (int)PPISPRegLossIndex::length> weights{};
for (auto& w : weights) w = r.uf(0.1f, 2.0f);
for (auto& w : weights) w = rr.uf(0.1f, 2.0f);
float* losses = alloc_zero<float>((int)PPISPRegLossIndex::length);
float* raw_losses = alloc_zero<float>((Bp + 1) * t.n_raw);
compute_ppsip_regularization_forward(
ttv(params, {Bp, t.n_params}), weights, t.name,
ttv(params, {Bp, t.n_params}), weights, t.name, arith != 0,
ttv(losses, {(int)PPISPRegLossIndex::length}),
ttv(raw_losses, {Bp + 1, t.n_raw}));
size_t raw_at = g_tight.size(), loss_at = g_loose.size() + t.n_raw;
readback_f(g_tight, raw_losses, Bp * t.n_raw);
readback_f(g_loose, raw_losses + Bp * t.n_raw, t.n_raw);
readback_f(g_loose, losses, (int)PPISPRegLossIndex::length);
// Backward with synthetic (host-fixed) summed raw losses so the
// whole chain stays deterministic.
std::vector<float> raw_h = r.vec((Bp + 1) * t.n_raw, -1.0f, 1.0f);
// whole chain stays deterministic. log2 needs a positive gain sum.
std::vector<float> raw_h = rr.vec((Bp + 1) * t.n_raw, -1.0f, 1.0f);
if (arith) raw_h[Bp * t.n_raw] = rr.uf(0.5f, 2.0f) * (float)Bp;
float* raw_fixed = upload(raw_h);
float* v_losses =
upload(r.vec((int)PPISPRegLossIndex::length, -1.0f, 1.0f));
std::vector<float> vl_h =
rr.vec((int)PPISPRegLossIndex::length, -1.0f, 1.0f);
float* v_losses = upload(vl_h);
float* v_params = alloc_zero<float>(Bp * t.n_params);
compute_ppsip_regularization_backward(
ttv(params, {Bp, t.n_params}), weights,
ttv(raw_fixed, {Bp + 1, t.n_raw}),
ttv(v_losses, {(int)PPISPRegLossIndex::length}), t.name,
ttv(v_params, {Bp, t.n_params}));
arith != 0, ttv(v_params, {Bp, t.n_params}));
size_t vp_at = g_tight.size();
readback_f(g_tight, v_params, Bp * t.n_params);
if (!arith) continue;
// Closed form of the exposure terms, which touch nothing else.
auto dsl1 = [](double x) {
return std::fabs(x) < 0.1 ? x / 0.1 : (x > 0 ? 1.0 : -1.0);
};
auto sl1 = [](double x) {
return std::fabs(x) < 0.1 ? 0.5 * x * x / 0.1 : std::fabs(x) - 0.05;
};
auto off = [](double got, double want) {
return std::fabs(got - want) > 1e-4 + 1e-4 * std::fabs(want);
};
double gain_sum = 0.0;
for (int64_t b = 0; b < Bp; b++) {
double gain = std::exp2((double)params_h[b * t.n_params]);
gain_sum += gain;
if (off(g_tight[raw_at + b * t.n_raw], gain)) {
std::fprintf(stderr, "ppisp_parity: %s raw gain[%lld] = %g, "
"expected %g\n", t.name, (long long)b,
g_tight[raw_at + b * t.n_raw], gain);
return 1;
}
}
double want_loss = weights[0] * sl1(std::log2(gain_sum / Bp));
if (off(g_loose[loss_at], want_loss)) {
std::fprintf(stderr, "ppisp_parity: %s exposure mean loss = %g, "
"expected %g\n", t.name, g_loose[loss_at], want_loss);
return 1;
}
double S = raw_h[Bp * t.n_raw];
double v_sum = vl_h[0] * weights[0] *
dsl1(std::log2(S / Bp)) / (S * std::log(2.0));
for (int64_t b = 0; b < Bp; b++) {
double p0 = params_h[b * t.n_params];
double want = v_sum * std::exp2(p0) * std::log(2.0);
if (off(g_tight[vp_at + b * t.n_params], want)) {
std::fprintf(stderr, "ppisp_parity: %s v_exposure[%lld] = %g, "
"expected %g\n", t.name, (long long)b,
g_tight[vp_at + b * t.n_params], want);
return 1;
}
}
}
}
+18 -11
View File
@@ -47,10 +47,13 @@ struct AddIntoGradParams {
};
static_assert(sizeof(AddIntoGradParams) == 2 * 8 + 2 * 4, "layout");
// (kParamType, kClampOutput), in ppisp_image.slang declaration order.
backend::vk::SpecList spec_list(const PpispParamSpec& spec) {
// (kParamType, kClampOutput, kExposureArithmeticMean), in ppisp_image.slang
// declaration order.
backend::vk::SpecList spec_list(const PpispParamSpec& spec,
bool exposure_arithmetic_mean) {
return backend::vk::SpecList{(uint32_t)spec.layout,
spec.clamp_output ? 1u : 0u};
spec.clamp_output ? 1u : 0u,
exposure_arithmetic_mean ? 1u : 0u};
}
// Fold the pixel axis across (gx, gz); gy carries the batch index.
@@ -62,8 +65,8 @@ void dispatch_image(const char* entry, const PpispParamSpec& spec,
p.wgs_per_row = f.per_row;
if (B > 65535 || f.rows > 65535)
throw std::runtime_error("ppisp: image grid dimension exceeds 65535");
vkk::dispatch(entry, spec_list(spec), f.per_row, (uint32_t)B, f.rows, &p,
sizeof(p));
vkk::dispatch(entry, spec_list(spec, false), f.per_row, (uint32_t)B,
f.rows, &p, sizeof(p));
}
} // namespace
@@ -130,11 +133,13 @@ void ppisp_backward(
void compute_ppsip_regularization_forward(
TorchTensorView ppisp_params,
const std::array<float, (int)PPISPRegLossIndex::length> loss_weights_0,
std::string param_type, TorchTensorView losses, TorchTensorView raw_losses
std::string param_type, bool exposure_arithmetic_mean,
TorchTensorView losses, TorchTensorView raw_losses
) {
const PpispParamSpec spec = ppisp_param_spec(param_type);
const int64_t B = std::get<2>(ppisp_params)[0];
const int nr = spec.num_raw_losses;
const backend::vk::SpecList sl = spec_list(spec, exposure_arithmetic_mean);
PpispRegParams p{};
p.ppisp_params = std::get<0>(ppisp_params);
@@ -144,7 +149,7 @@ void compute_ppsip_regularization_forward(
for (int i = 0; i < (int)PPISPRegLossIndex::length; i++)
p.loss_weights[i] = loss_weights_0[i];
p.B = (int32_t)B;
vkk::dispatch("ppisp_image.ppisp_reg_raw_fwd", spec_list(spec),
vkk::dispatch("ppisp_image.ppisp_reg_raw_fwd", sl,
(uint32_t)((B + 31) / 32), 1, 1, &p, sizeof(p));
// Weighted final pass reads the summed tail row.
@@ -157,7 +162,7 @@ void compute_ppsip_regularization_forward(
for (int i = 0; i < (int)PPISPRegLossIndex::length; i++)
pf.loss_weights[i] = loss_weights_0[i];
pf.B = (int32_t)B;
vkk::dispatch("ppisp_image.ppisp_reg_final_fwd", spec_list(spec), 1, 1, 1,
vkk::dispatch("ppisp_image.ppisp_reg_final_fwd", sl, 1, 1, 1,
&pf, sizeof(pf));
}
@@ -165,11 +170,13 @@ void compute_ppsip_regularization_backward(
TorchTensorView ppisp_params,
const std::array<float, (int)PPISPRegLossIndex::length> loss_weights_0,
TorchTensorView raw_losses, TorchTensorView v_losses,
std::string param_type, TorchTensorView v_ppisp_params
std::string param_type, bool exposure_arithmetic_mean,
TorchTensorView v_ppisp_params
) {
const PpispParamSpec spec = ppisp_param_spec(param_type);
const int64_t B = std::get<2>(ppisp_params)[0];
const int nr = spec.num_raw_losses;
const backend::vk::SpecList sl = spec_list(spec, exposure_arithmetic_mean);
float* v_raw_losses =
DevicePool::global().acquire<float>(PoolSlot::PpispVRawLosses, nr);
@@ -185,7 +192,7 @@ void compute_ppsip_regularization_backward(
for (int i = 0; i < (int)PPISPRegLossIndex::length; i++)
pf.loss_weights[i] = loss_weights_0[i];
pf.B = (int32_t)B;
vkk::dispatch("ppisp_image.ppisp_reg_final_bwd", spec_list(spec), 1, 1, 1,
vkk::dispatch("ppisp_image.ppisp_reg_final_bwd", sl, 1, 1, 1,
&pf, sizeof(pf));
// v_raw_losses -> v_ppisp_params (per image).
@@ -197,7 +204,7 @@ void compute_ppsip_regularization_backward(
for (int i = 0; i < (int)PPISPRegLossIndex::length; i++)
p.loss_weights[i] = loss_weights_0[i];
p.B = (int32_t)B;
vkk::dispatch("ppisp_image.ppisp_reg_raw_bwd", spec_list(spec),
vkk::dispatch("ppisp_image.ppisp_reg_raw_bwd", sl,
(uint32_t)((B + 31) / 32), 1, 1, &p, sizeof(p));
}
+38 -20
View File
@@ -1,13 +1,12 @@
// PPISP image transform + regularization losses (Vulkan backend), mirroring
// the PPISP kernels in kernels/ppisp/Ppisp.cu. All device math comes from the
// canonical shaders/ppisp.slang (same functions the CUDA build calls through
// generated/ppisp.cuh), so forward values match up to fast-math rounding and
// gradients are the same autodiff.
// kernels/ppisp/Ppisp.cu. All device math is shaders/ppisp.slang, the same
// functions the CUDA build calls through generated/ppisp.cuh.
//
// Spec constants (IDs by declaration order — dispatches must pass ALL):
// 0: kParamType — 0 = Original (36), 1 = RQS (39), 2 = NoCRF (24),
// 3 = NoCRFNoVig (9)
// 1: kClampOutput — clamp the result to [0,1] (CRF-less layouts only)
// 2: kExposureArithmeticMean — regularize log2(mean gain), not mean(log2 gain)
#include "atomic_float.slang"
@@ -21,6 +20,9 @@ const int kParamType = 0;
[SpecializationConstant]
const int kClampOutput = 0;
[SpecializationConstant]
const int kExposureArithmeticMean = 0;
static const int kMaxParams = PPISP_RQS_NUM_PARAMS; // 39
static const int kMaxRawLosses = int(RawPPISPRegLossIndexRQS::length); // 23
static const int kNumWeighted = int(PPISPRegLossIndex::length); // 6
@@ -253,7 +255,8 @@ void _pp_raw_losses(float pbuf[kMaxParams], out float lbuf[kMaxRawLosses]) {
[ForceUnroll]
for (int i = 0; i < PPISP_NUM_PARAMS; i++) ps[i] = pbuf[i];
float ls[int(RawPPISPRegLossIndex::length)] =
compute_raw_ppisp_regularization_loss(ps);
compute_raw_ppisp_regularization_loss(
ps, kExposureArithmeticMean != 0);
[ForceUnroll]
for (int i = 0; i < int(RawPPISPRegLossIndex::length); i++)
lbuf[i] = ls[i];
@@ -262,7 +265,8 @@ void _pp_raw_losses(float pbuf[kMaxParams], out float lbuf[kMaxRawLosses]) {
[ForceUnroll]
for (int i = 0; i < PPISP_RQS_NUM_PARAMS; i++) ps[i] = pbuf[i];
float ls[int(RawPPISPRegLossIndexRQS::length)] =
compute_raw_ppisp_rqs_regularization_loss(ps);
compute_raw_ppisp_rqs_regularization_loss(
ps, kExposureArithmeticMean != 0);
[ForceUnroll]
for (int i = 0; i < int(RawPPISPRegLossIndexRQS::length); i++)
lbuf[i] = ls[i];
@@ -271,7 +275,8 @@ void _pp_raw_losses(float pbuf[kMaxParams], out float lbuf[kMaxRawLosses]) {
[ForceUnroll]
for (int i = 0; i < PPISP_NO_CRF_NUM_PARAMS; i++) ps[i] = pbuf[i];
float ls[int(RawPPISPRegLossIndexNoCRF::length)] =
compute_raw_ppisp_no_crf_regularization_loss(ps);
compute_raw_ppisp_no_crf_regularization_loss(
ps, kExposureArithmeticMean != 0);
[ForceUnroll]
for (int i = 0; i < int(RawPPISPRegLossIndexNoCRF::length); i++)
lbuf[i] = ls[i];
@@ -280,7 +285,8 @@ void _pp_raw_losses(float pbuf[kMaxParams], out float lbuf[kMaxRawLosses]) {
[ForceUnroll]
for (int i = 0; i < PPISP_NO_CRF_NO_VIG_NUM_PARAMS; i++) ps[i] = pbuf[i];
float ls[int(RawPPISPRegLossIndexNoCRFNoVig::length)] =
compute_raw_ppisp_no_crf_no_vig_regularization_loss(ps);
compute_raw_ppisp_no_crf_no_vig_regularization_loss(
ps, kExposureArithmeticMean != 0);
[ForceUnroll]
for (int i = 0; i < int(RawPPISPRegLossIndexNoCRFNoVig::length); i++)
lbuf[i] = ls[i];
@@ -348,7 +354,8 @@ void ppisp_reg_final_fwd(uint tid: SV_GroupThreadID,
for (int i = 0; i < int(RawPPISPRegLossIndex::length); i++)
raw[i] = p.raw_losses[i];
float o[kNumWeighted] =
compute_ppisp_regularization_loss(raw, p.B, w);
compute_ppisp_regularization_loss(
raw, p.B, kExposureArithmeticMean != 0, w);
[ForceUnroll]
for (int i = 0; i < kNumWeighted; i++) ls[i] = o[i];
} else if (kParamType == 1) {
@@ -357,7 +364,8 @@ void ppisp_reg_final_fwd(uint tid: SV_GroupThreadID,
for (int i = 0; i < int(RawPPISPRegLossIndexRQS::length); i++)
raw[i] = p.raw_losses[i];
float o[kNumWeighted] =
compute_ppisp_rqs_regularization_loss(raw, p.B, w);
compute_ppisp_rqs_regularization_loss(
raw, p.B, kExposureArithmeticMean != 0, w);
[ForceUnroll]
for (int i = 0; i < kNumWeighted; i++) ls[i] = o[i];
} else if (kParamType == 2) {
@@ -366,7 +374,8 @@ void ppisp_reg_final_fwd(uint tid: SV_GroupThreadID,
for (int i = 0; i < int(RawPPISPRegLossIndexNoCRF::length); i++)
raw[i] = p.raw_losses[i];
float o[kNumWeighted] =
compute_ppisp_no_crf_regularization_loss(raw, p.B, w);
compute_ppisp_no_crf_regularization_loss(
raw, p.B, kExposureArithmeticMean != 0, w);
[ForceUnroll]
for (int i = 0; i < kNumWeighted; i++) ls[i] = o[i];
} else {
@@ -375,7 +384,8 @@ void ppisp_reg_final_fwd(uint tid: SV_GroupThreadID,
for (int i = 0; i < int(RawPPISPRegLossIndexNoCRFNoVig::length); i++)
raw[i] = p.raw_losses[i];
float o[kNumWeighted] =
compute_ppisp_no_crf_no_vig_regularization_loss(raw, p.B, w);
compute_ppisp_no_crf_no_vig_regularization_loss(
raw, p.B, kExposureArithmeticMean != 0, w);
[ForceUnroll]
for (int i = 0; i < kNumWeighted; i++) ls[i] = o[i];
}
@@ -405,7 +415,8 @@ void ppisp_reg_final_bwd(uint tid: SV_GroupThreadID,
for (int i = 0; i < int(RawPPISPRegLossIndex::length); i++)
raw[i] = p.raw_losses[i];
float o[int(RawPPISPRegLossIndex::length)] =
compute_ppisp_regularization_loss_vjp(raw, p.B, w, vl);
compute_ppisp_regularization_loss_vjp(
raw, p.B, kExposureArithmeticMean != 0, w, vl);
[ForceUnroll]
for (int i = 0; i < int(RawPPISPRegLossIndex::length); i++)
p.v_ppisp_params[i] = o[i];
@@ -415,7 +426,8 @@ void ppisp_reg_final_bwd(uint tid: SV_GroupThreadID,
for (int i = 0; i < int(RawPPISPRegLossIndexRQS::length); i++)
raw[i] = p.raw_losses[i];
float o[int(RawPPISPRegLossIndexRQS::length)] =
compute_ppisp_rqs_regularization_loss_vjp(raw, p.B, w, vl);
compute_ppisp_rqs_regularization_loss_vjp(
raw, p.B, kExposureArithmeticMean != 0, w, vl);
[ForceUnroll]
for (int i = 0; i < int(RawPPISPRegLossIndexRQS::length); i++)
p.v_ppisp_params[i] = o[i];
@@ -425,7 +437,8 @@ void ppisp_reg_final_bwd(uint tid: SV_GroupThreadID,
for (int i = 0; i < int(RawPPISPRegLossIndexNoCRF::length); i++)
raw[i] = p.raw_losses[i];
float o[int(RawPPISPRegLossIndexNoCRF::length)] =
compute_ppisp_no_crf_regularization_loss_vjp(raw, p.B, w, vl);
compute_ppisp_no_crf_regularization_loss_vjp(
raw, p.B, kExposureArithmeticMean != 0, w, vl);
[ForceUnroll]
for (int i = 0; i < int(RawPPISPRegLossIndexNoCRF::length); i++)
p.v_ppisp_params[i] = o[i];
@@ -435,7 +448,8 @@ void ppisp_reg_final_bwd(uint tid: SV_GroupThreadID,
for (int i = 0; i < int(RawPPISPRegLossIndexNoCRFNoVig::length); i++)
raw[i] = p.raw_losses[i];
float o[int(RawPPISPRegLossIndexNoCRFNoVig::length)] =
compute_ppisp_no_crf_no_vig_regularization_loss_vjp(raw, p.B, w, vl);
compute_ppisp_no_crf_no_vig_regularization_loss_vjp(
raw, p.B, kExposureArithmeticMean != 0, w, vl);
[ForceUnroll]
for (int i = 0; i < int(RawPPISPRegLossIndexNoCRFNoVig::length); i++)
p.v_ppisp_params[i] = o[i];
@@ -468,7 +482,8 @@ void ppisp_reg_raw_bwd(uint3 wg: SV_GroupID, uint tid: SV_GroupThreadID,
for (int i = 0; i < int(RawPPISPRegLossIndex::length); i++)
vr[i] = p.raw_losses[i];
float o[PPISP_NUM_PARAMS] =
compute_raw_ppisp_regularization_loss_vjp(ps, vr);
compute_raw_ppisp_regularization_loss_vjp(
ps, kExposureArithmeticMean != 0, vr);
[ForceUnroll]
for (int i = 0; i < PPISP_NUM_PARAMS; i++) vbuf[i] = o[i];
} else if (kParamType == 1) {
@@ -480,7 +495,8 @@ void ppisp_reg_raw_bwd(uint3 wg: SV_GroupID, uint tid: SV_GroupThreadID,
for (int i = 0; i < int(RawPPISPRegLossIndexRQS::length); i++)
vr[i] = p.raw_losses[i];
float o[PPISP_RQS_NUM_PARAMS] =
compute_raw_ppisp_rqs_regularization_loss_vjp(ps, vr);
compute_raw_ppisp_rqs_regularization_loss_vjp(
ps, kExposureArithmeticMean != 0, vr);
[ForceUnroll]
for (int i = 0; i < PPISP_RQS_NUM_PARAMS; i++) vbuf[i] = o[i];
} else if (kParamType == 2) {
@@ -492,7 +508,8 @@ void ppisp_reg_raw_bwd(uint3 wg: SV_GroupID, uint tid: SV_GroupThreadID,
for (int i = 0; i < int(RawPPISPRegLossIndexNoCRF::length); i++)
vr[i] = p.raw_losses[i];
float o[PPISP_NO_CRF_NUM_PARAMS] =
compute_raw_ppisp_no_crf_regularization_loss_vjp(ps, vr);
compute_raw_ppisp_no_crf_regularization_loss_vjp(
ps, kExposureArithmeticMean != 0, vr);
[ForceUnroll]
for (int i = 0; i < PPISP_NO_CRF_NUM_PARAMS; i++) vbuf[i] = o[i];
} else {
@@ -504,7 +521,8 @@ void ppisp_reg_raw_bwd(uint3 wg: SV_GroupID, uint tid: SV_GroupThreadID,
for (int i = 0; i < int(RawPPISPRegLossIndexNoCRFNoVig::length); i++)
vr[i] = p.raw_losses[i];
float o[PPISP_NO_CRF_NO_VIG_NUM_PARAMS] =
compute_raw_ppisp_no_crf_no_vig_regularization_loss_vjp(ps, vr);
compute_raw_ppisp_no_crf_no_vig_regularization_loss_vjp(
ps, kExposureArithmeticMean != 0, vr);
[ForceUnroll]
for (int i = 0; i < PPISP_NO_CRF_NO_VIG_NUM_PARAMS; i++) vbuf[i] = o[i];
}
+2 -1
View File
@@ -258,6 +258,7 @@ inline int train_tier_rank(const char* tier) {
X(bool, use_ppisp, true, "correction", "basic", "") \
X(std::string, ppisp_param_type, "no_crf_no_vig", "correction", "basic", "original|rqs|no_crf|no_crf_clamp|no_crf_no_vig|no_crf_no_vig_clamp") \
X(bool, ppisp_exposure_from_exif, false, "correction", "basic", "") \
X(bool, ppisp_exposure_arithmetic_mean, true, "correction", "expert", "") \
X(bool, apply_ppisp_before_bilagrid, true, "correction", "advanced", "") \
X(bool, apply_ppisp_before_color_space, false, "correction", "advanced", "") \
X(bool, use_adagrad_ppisp_optim, true, "correction", "advanced", "") \
@@ -431,7 +432,7 @@ inline bool train_apply_preset(TrainConfig& c, const std::string& name) {
// c.apply_ppisp_before_color_space = true;
// c.ppisp_adagrad_lr = 0.25f;
c.ppisp_exposure_from_exif = true;
// c.background_mode = "noise";
c.background_mode = "random";
// c.depth_distortion_reg = 0.01f;
c.loss_saturation_threshold = 0.98f;
c.normalize_loss_by_luminance = true;
+4 -1
View File
@@ -231,8 +231,11 @@ void engine_bilagrid_optim_step(int step, const BilagridStepConfig& cfg);
// --- PPISP (RGB only; PpispStepConfig picks where in the chain it runs) ---
// ppisp_param_spec (kernels/pixelwise/PixelWise.cuh) owns the param_type list.
// exposure_init optionally seeds params[:, 0] with [n_grids] log2 gains.
// exposure_init optionally seeds params[:, 0] with [n_grids] log2 gains;
// exposure_arithmetic_mean regularizes log2(mean gain), else mean(log2 gain).
void engine_init_ppisp(int n_grids, std::string param_type, bool use_adagrad,
bool exposure_arithmetic_mean,
const std::vector<float>& exposure_init = {});
// Apply PPISP forward in place on the current rendered RGB; saves a pre-PPISP
+5 -2
View File
@@ -39,6 +39,7 @@ static TorchTensorView _ppisp_cam_indices_tv() {
}
void engine_init_ppisp(int n_grids, std::string param_type, bool use_adagrad,
bool exposure_arithmetic_mean,
const std::vector<float>& exposure_init) {
if (n_grids <= 0)
throw std::runtime_error("engine_init_ppisp: n_grids must be > 0");
@@ -50,6 +51,7 @@ void engine_init_ppisp(int n_grids, std::string param_type, bool use_adagrad,
engine().ppisp.param_type = (param_type == "" ? std::string("original") : param_type);
engine().ppisp.num_params = P;
engine().ppisp.use_adagrad = use_adagrad;
engine().ppisp.exposure_arithmetic_mean = exposure_arithmetic_mean;
engine().ppisp.params.resize(PoolSlot::EngPpispParams, n_grids, P);
if (engine().ppisp.param_type == "original") {
ppisp_original_default_init(
@@ -193,7 +195,7 @@ float* _engine_ppisp_reg_loss_into(
compute_ppsip_regularization_forward(
params_tv, loss_weights, engine().ppisp.param_type,
losses_tv, raw_tv);
engine().ppisp.exposure_arithmetic_mean, losses_tv, raw_tv);
if (compute_grad) {
// v_losses = ones[kLoss]: gradient flows back through reg-loss sum.
@@ -217,7 +219,8 @@ float* _engine_ppisp_reg_loss_into(
compute_ppsip_regularization_backward(
params_tv, loss_weights, raw_tv, v_losses_tv,
engine().ppisp.param_type, v_params_tv);
engine().ppisp.param_type, engine().ppisp.exposure_arithmetic_mean,
v_params_tv);
// ppisp_grads += v_params_scratch (over all N * P floats).
size_t total = (size_t)N * engine().ppisp.num_params;
+1
View File
@@ -461,6 +461,7 @@ struct PpispState {
bool enabled = false;
bool optim_initialized = false;
bool use_adagrad = false;
bool exposure_arithmetic_mean = false;
// Per-iteration mirror of the PpispStepConfig order flags, stashed by the
// forward path so the backward hooks in EngineLoss.cpp can invert the
// order they picked. Reset each step.
Generated Vendored
+7781 -7250
View File
File diff suppressed because it is too large Load Diff
+74
View File
@@ -7606,6 +7606,80 @@ SS_MSG(ppisp_exposure_from_exif_help,
"etiketleri olmayan fotoğraflar ortalamadan başlar. Pozlama çekim boyunca "
"değişiyorsa yardımcı olur."));
SS_MSG(ppisp_exposure_arithmetic_mean,
EN("Neutral exposure by average gain"), JA("平均の倍率で露出を中立に"),
ZH_HANS("按平均倍率保持曝光中性"), ZH_HANT("按平均倍率保持曝光中性"),
KO("평균 배율로 노출 중립"), DE("Neutrale Belichtung über mittleren Faktor"),
FR("Exposition neutre par gain moyen"),
ES("Exposición neutra por ganancia media"),
PT("Exposição neutra pelo ganho médio"),
IT("Esposizione neutra per guadagno medio"),
NL("Neutrale belichting via gemiddelde factor"),
RU("Нейтральная экспозиция по среднему множителю"),
TR("Ortalama çarpanla nötr pozlama"));
SS_MSG(ppisp_exposure_arithmetic_mean_help,
EN("Center the per-photo exposure corrections so their brightness multipliers "
"average to 1, rather than their values in stops averaging to 0. Applies "
"to both the neutral exposure penalty and the EXIF start. When exposure "
"varies widely, this keeps the splats at the photos' average brightness "
"instead of darker."),
JA("写真ごとの露出補正を、段数での値の平均が 0 になるようにではなく、明るさ"
"の倍率の平均が 1 になるように中心を合わせます。露出を中立に保つ強さと "
"EXIF による露出の初期化の両方に適用されます。露出の差が大きいとき、スプ"
"ラットが暗くならず、写真の平均的な明るさに保たれます。"),
ZH_HANS("让逐张照片的曝光校正以亮度倍率的平均值为 1 为中心,而不是以档数"
"的平均值为 0。同时作用于保持曝光中性的强度和用 EXIF 初始化曝光。"
"曝光差异很大时,这会让泼溅保持在照片的平均亮度,而不是更暗。"),
ZH_HANT("讓逐張照片的曝光校正以亮度倍率的平均值為 1 為中心,而不是以檔數"
"的平均值為 0。同時作用於保持曝光中性的強度和用 EXIF 初始化曝光。"
"曝光差異很大時,這會讓潑濺保持在照片的平均亮度,而不是更暗。"),
KO("사진별 노출 보정을, 스톱 단위 값의 평균이 0이 되도록이 아니라 밝기 배율"
"의 평균이 1이 되도록 맞춥니다. 노출을 중립으로 유지하는 강도와 EXIF로 노"
"출 초기화에 모두 적용됩니다. 노출 차이가 클 때 스플랫이 더 어두워지지 않"
"고 사진의 평균 밝기에 맞춰집니다."),
DE("Die Belichtungskorrekturen pro Foto so zentrieren, dass ihre "
"Helligkeitsfaktoren im Mittel 1 ergeben, statt dass ihre Werte in "
"Blendenstufen im Mittel 0 ergeben. Gilt für die Strafe für nicht neutrale "
"Belichtung und den Belichtungsstart aus EXIF. Bei stark schwankender "
"Belichtung bleiben die Splats so bei der mittleren Helligkeit der Fotos "
"statt dunkler."),
FR("Centrer les corrections d'exposition par photo pour que leurs facteurs de "
"luminosité aient une moyenne de 1, plutôt que leurs valeurs en stops une "
"moyenne de 0. S'applique à la pénalité d'exposition non neutre et à "
"l'exposition initiale depuis l'EXIF. Quand l'exposition varie fortement, "
"les splats restent ainsi à la luminosité moyenne des photos au lieu "
"d'être plus sombres."),
ES("Centrar las correcciones de exposición por foto para que sus factores de "
"brillo promedien 1, en lugar de que sus valores en pasos promedien 0. Se "
"aplica a la penalización de exposición no neutra y a la exposición inicial "
"desde EXIF. Cuando la exposición varía mucho, así los splats quedan con el "
"brillo medio de las fotos en vez de más oscuros."),
PT("Centralizar as correções de exposição por foto para que seus fatores de "
"brilho tenham média 1, em vez de seus valores em stops terem média 0. Vale "
"para a penalidade de exposição não neutra e para a exposição inicial do "
"EXIF. Quando a exposição varia muito, os splats ficam assim no brilho "
"médio das fotos em vez de mais escuros."),
IT("Centrare le correzioni di esposizione di ogni foto in modo che i loro "
"fattori di luminosità abbiano media 1, invece che i loro valori in stop "
"abbiano media 0. Vale per la penalità di esposizione non neutra e per "
"l'esposizione iniziale da EXIF. Quando l'esposizione varia molto, gli "
"splat restano così alla luminosità media delle foto invece che più scuri."),
NL("De belichtingscorrecties per foto zo centreren dat hun helderheidsfactoren "
"gemiddeld 1 zijn, in plaats van dat hun waarden in stops gemiddeld 0 zijn. "
"Geldt voor de straf voor niet-neutrale belichting en voor het starten van "
"de belichting vanuit EXIF. Bij sterk wisselende belichting blijven de "
"splats zo op de gemiddelde helderheid van de foto's in plaats van donkerder."),
RU("Центрировать коррекции экспозиции каждого фото так, чтобы в среднем 1 "
"давали их множители яркости, а не 0 — их значения в ступенях. Действует и "
"на штраф за смещение экспозиции, и на начальную экспозицию из EXIF. При "
"сильно различающейся экспозиции сплаты так остаются на средней яркости "
"фотографий, а не темнее."),
TR("Fotoğraf başına pozlama düzeltmelerini, durak cinsinden değerlerinin "
"ortalaması 0 olacak şekilde değil, parlaklık çarpanlarının ortalaması 1 "
"olacak şekilde ortalar. Hem nötr olmayan pozlama cezasına hem de EXIF'ten "
"pozlama başlangıcına uygulanır. Pozlama çok değiştiğinde splat'lar böylece "
"daha karanlık kalmak yerine fotoğrafların ortalama parlaklığında kalır."));
SS_MSG(apply_ppisp_before_bilagrid,
EN("Camera correction first"), JA("カメラ補正を先に適用"),
ZH_HANS("先做相机校正"), ZH_HANT("先做相機校正"),
+2
View File
@@ -595,6 +595,7 @@ void compute_ppsip_regularization_forward(
TorchTensorView ppisp_params, // [B, PPISP_NUM_PARAMS]
const std::array<float, (int)PPISPRegLossIndex::length> loss_weights_0,
std::string param_type,
bool exposure_arithmetic_mean, // log2(mean gain) = 0, else mean(log2 gain) = 0
TorchTensorView losses, // [PPISPRegLossIndex::length] (must be pre-zeroed)
TorchTensorView raw_losses // [B+1, RawPPISPRegLossIndex::length] (must be pre-zeroed)
);
@@ -606,5 +607,6 @@ void compute_ppsip_regularization_backward(
TorchTensorView raw_losses, // [B+1, RawPPISPRegLossIndex::length]
TorchTensorView v_losses, // [PPISPRegLossIndex::length]
std::string param_type,
bool exposure_arithmetic_mean,
TorchTensorView v_ppisp_params // [B, PPISP_NUM_PARAMS] (must be pre-zeroed)
);
+50 -26
View File
@@ -48,6 +48,17 @@ static void _with_ppisp_layout(PpispParamLayout layout, F&& f) {
}
}
template<class F>
static void _with_ppisp_reg_mode(PpispParamLayout layout,
bool exposure_arithmetic_mean, F&& f) {
_with_ppisp_layout(layout, [&](auto L) {
if (exposure_arithmetic_mean)
f(L, std::true_type{});
else
f(L, std::false_type{});
});
}
template<PpispParamLayout layout>
__global__ void ppisp_forward_kernel(
const TensorView<float, 4> in_image, // [B, H, W, C]
@@ -256,7 +267,7 @@ void ppisp_backward(
CHECK_DEVICE_ERROR(cudaGetLastError());
}
template<PpispParamLayout layout>
template<PpispParamLayout layout, bool exposure_arithmetic_mean>
__global__ void compute_raw_ppisp_regularization_forward_kernel(
int B, // number of images
const float* __restrict__ ppisp_params, // [B, PPISP_NUM_PARAMS]
@@ -277,14 +288,17 @@ __global__ void compute_raw_ppisp_regularization_forward_kernel(
}
if constexpr (layout == PpispParamLayout::Original)
SlangPPISP::compute_raw_ppisp_regularization_loss(params, &losses);
SlangPPISP::compute_raw_ppisp_regularization_loss(
params, exposure_arithmetic_mean, &losses);
else if constexpr (layout == PpispParamLayout::RQS)
SlangPPISP::compute_raw_ppisp_rqs_regularization_loss(params, &losses);
SlangPPISP::compute_raw_ppisp_rqs_regularization_loss(
params, exposure_arithmetic_mean, &losses);
else if constexpr (layout == PpispParamLayout::NoCRF)
SlangPPISP::compute_raw_ppisp_no_crf_regularization_loss(params, &losses);
SlangPPISP::compute_raw_ppisp_no_crf_regularization_loss(
params, exposure_arithmetic_mean, &losses);
else
SlangPPISP::compute_raw_ppisp_no_crf_no_vig_regularization_loss(
params, &losses);
params, exposure_arithmetic_mean, &losses);
}
auto block = cg::this_thread_block();
@@ -300,7 +314,7 @@ __global__ void compute_raw_ppisp_regularization_forward_kernel(
}
}
template<PpispParamLayout layout>
template<PpispParamLayout layout, bool exposure_arithmetic_mean>
__global__ void compute_ppisp_regularization_forward_kernel(
int num_train_images,
const float* __restrict__ raw_losses_buffer, // [RawPPISPRegLossIndex::length]
@@ -319,16 +333,20 @@ __global__ void compute_ppisp_regularization_forward_kernel(
if constexpr (layout == PpispParamLayout::Original)
SlangPPISP::compute_ppisp_regularization_loss(
raw_losses, num_train_images, loss_weights, &losses);
raw_losses, num_train_images, exposure_arithmetic_mean, loss_weights,
&losses);
else if constexpr (layout == PpispParamLayout::RQS)
SlangPPISP::compute_ppisp_rqs_regularization_loss(
raw_losses, num_train_images, loss_weights, &losses);
raw_losses, num_train_images, exposure_arithmetic_mean, loss_weights,
&losses);
else if constexpr (layout == PpispParamLayout::NoCRF)
SlangPPISP::compute_ppisp_no_crf_regularization_loss(
raw_losses, num_train_images, loss_weights, &losses);
raw_losses, num_train_images, exposure_arithmetic_mean, loss_weights,
&losses);
else
SlangPPISP::compute_ppisp_no_crf_no_vig_regularization_loss(
raw_losses, num_train_images, loss_weights, &losses);
raw_losses, num_train_images, exposure_arithmetic_mean, loss_weights,
&losses);
#pragma unroll
for (int i = 0; i < (int)PPISPRegLossIndex::length; i++) {
@@ -341,6 +359,7 @@ void compute_ppsip_regularization_forward(
TorchTensorView ppisp_params, // [B, PPISP_NUM_PARAMS]
const std::array<float, (int)PPISPRegLossIndex::length> loss_weights_0,
std::string param_type,
bool exposure_arithmetic_mean, // log2(mean gain) = 0, else mean(log2 gain) = 0
TorchTensorView losses, // [PPISPRegLossIndex::length] (must be pre-zeroed)
TorchTensorView raw_losses // [B+1, RawPPISPRegLossIndex::length] (must be pre-zeroed)
) {
@@ -350,8 +369,8 @@ void compute_ppsip_regularization_forward(
long B = std::get<2>(ppisp_params)[0];
const PpispParamSpec spec = ppisp_param_spec(param_type);
_with_ppisp_layout(spec.layout, [&](auto L) {
compute_raw_ppisp_regularization_forward_kernel<L.value>
_with_ppisp_reg_mode(spec.layout, exposure_arithmetic_mean, [&](auto L, auto A) {
compute_raw_ppisp_regularization_forward_kernel<L.value, A.value>
<<<_LAUNCH_ARGS_1D(B, WARP_SIZE)>>>(
B,
(float*)std::get<0>(ppisp_params),
@@ -359,7 +378,7 @@ void compute_ppsip_regularization_forward(
);
CHECK_DEVICE_ERROR(cudaGetLastError());
compute_ppisp_regularization_forward_kernel<L.value>
compute_ppisp_regularization_forward_kernel<L.value, A.value>
<<<1, 1>>>(
B,
(float*)std::get<0>(raw_losses) + B * spec.num_raw_losses,
@@ -370,7 +389,7 @@ void compute_ppsip_regularization_forward(
});
}
template<PpispParamLayout layout>
template<PpispParamLayout layout, bool exposure_arithmetic_mean>
__global__ void compute_raw_ppisp_regularization_backward_kernel(
int B, // number of images
const float* __restrict__ ppisp_params, // [B, PPISP_NUM_PARAMS]
@@ -399,16 +418,16 @@ __global__ void compute_raw_ppisp_regularization_backward_kernel(
FixedArray<float, kNumParams> v_params;
if constexpr (layout == PpispParamLayout::Original)
SlangPPISP::compute_raw_ppisp_regularization_loss_vjp(
params, v_losses, &v_params);
params, exposure_arithmetic_mean, v_losses, &v_params);
else if constexpr (layout == PpispParamLayout::RQS)
SlangPPISP::compute_raw_ppisp_rqs_regularization_loss_vjp(
params, v_losses, &v_params);
params, exposure_arithmetic_mean, v_losses, &v_params);
else if constexpr (layout == PpispParamLayout::NoCRF)
SlangPPISP::compute_raw_ppisp_no_crf_regularization_loss_vjp(
params, v_losses, &v_params);
params, exposure_arithmetic_mean, v_losses, &v_params);
else
SlangPPISP::compute_raw_ppisp_no_crf_no_vig_regularization_loss_vjp(
params, v_losses, &v_params);
params, exposure_arithmetic_mean, v_losses, &v_params);
#pragma unroll
for (int i = 0; i < kNumParams; i++) {
@@ -417,7 +436,7 @@ __global__ void compute_raw_ppisp_regularization_backward_kernel(
}
}
template<PpispParamLayout layout>
template<PpispParamLayout layout, bool exposure_arithmetic_mean>
__global__ void compute_ppisp_regularization_backward_kernel(
int num_train_images,
const float* __restrict__ raw_losses_buffer, // [RawPPISPRegLossIndex::length]
@@ -442,16 +461,20 @@ __global__ void compute_ppisp_regularization_backward_kernel(
FixedArray<float, kNumRawLosses> v_raw_losses;
if constexpr (layout == PpispParamLayout::Original)
SlangPPISP::compute_ppisp_regularization_loss_vjp(
raw_losses, num_train_images, loss_weights, v_losses, &v_raw_losses);
raw_losses, num_train_images, exposure_arithmetic_mean, loss_weights,
v_losses, &v_raw_losses);
else if constexpr (layout == PpispParamLayout::RQS)
SlangPPISP::compute_ppisp_rqs_regularization_loss_vjp(
raw_losses, num_train_images, loss_weights, v_losses, &v_raw_losses);
raw_losses, num_train_images, exposure_arithmetic_mean, loss_weights,
v_losses, &v_raw_losses);
else if constexpr (layout == PpispParamLayout::NoCRF)
SlangPPISP::compute_ppisp_no_crf_regularization_loss_vjp(
raw_losses, num_train_images, loss_weights, v_losses, &v_raw_losses);
raw_losses, num_train_images, exposure_arithmetic_mean, loss_weights,
v_losses, &v_raw_losses);
else
SlangPPISP::compute_ppisp_no_crf_no_vig_regularization_loss_vjp(
raw_losses, num_train_images, loss_weights, v_losses, &v_raw_losses);
raw_losses, num_train_images, exposure_arithmetic_mean, loss_weights,
v_losses, &v_raw_losses);
#pragma unroll
for (int i = 0; i < kNumRawLosses; i++) {
@@ -466,6 +489,7 @@ void compute_ppsip_regularization_backward(
TorchTensorView raw_losses, // [B+1, RawPPISPRegLossIndex::length]
TorchTensorView v_losses, // [PPISPRegLossIndex::length]
std::string param_type,
bool exposure_arithmetic_mean,
TorchTensorView v_ppisp_params // [B, PPISP_NUM_PARAMS] (must be pre-zeroed)
) {
FixedArray<float, (int)PPISPRegLossIndex::length> loss_weights =
@@ -479,8 +503,8 @@ void compute_ppsip_regularization_backward(
PoolSlot::PpispVRawLosses, spec.num_raw_losses);
cudaMemset(v_raw_losses, 0, spec.num_raw_losses * sizeof(float));
_with_ppisp_layout(spec.layout, [&](auto L) {
compute_ppisp_regularization_backward_kernel<L.value>
_with_ppisp_reg_mode(spec.layout, exposure_arithmetic_mean, [&](auto L, auto A) {
compute_ppisp_regularization_backward_kernel<L.value, A.value>
<<<1, 1>>>(
B,
(float*)std::get<0>(raw_losses) + B * spec.num_raw_losses,
@@ -490,7 +514,7 @@ void compute_ppsip_regularization_backward(
);
CHECK_DEVICE_ERROR(cudaGetLastError());
compute_raw_ppisp_regularization_backward_kernel<L.value>
compute_raw_ppisp_regularization_backward_kernel<L.value, A.value>
<<<_LAUNCH_ARGS_1D(B, WARP_SIZE)>>>(
B,
(float*)std::get<0>(ppisp_params),
+65 -28
View File
@@ -711,6 +711,16 @@ enum RawPPISPRegLossIndexRQS {
length
};
// Summed over cameras: log2(E[gain]) = 0 needs the gains, E[log2 gain] = 0
// the exposure parameters themselves.
[ForceInline]
[Differentiable]
float ppisp_raw_exposure(float exposure, bool arithmetic_mean) {
if (arithmetic_mean)
return exp2(exposure);
return exposure;
}
[ForceInline]
[Differentiable]
float varp_of_3(float a, float b, float c) {
@@ -723,13 +733,14 @@ float varp_of_3(float a, float b, float c) {
[CudaDeviceExport]
[Differentiable]
float[RawPPISPRegLossIndex::length] compute_raw_ppisp_regularization_loss(
float[PPISP_NUM_PARAMS] params
float[PPISP_NUM_PARAMS] params, bool exposure_arithmetic_mean
) {
PPISPParams p = get_ppisp_params(params);
float losses[RawPPISPRegLossIndex::length] = {0.0f};
// Exposure
losses[RawPPISPRegLossIndex::SumExposure] = p.exposure;
losses[RawPPISPRegLossIndex::SumExposure] =
ppisp_raw_exposure(p.exposure, exposure_arithmetic_mean);
// Vignetting
losses[RawPPISPRegLossIndex::SumVignettingCrSquared] =
@@ -793,12 +804,13 @@ float[RawPPISPRegLossIndex::length] compute_raw_ppisp_regularization_loss(
[CudaDeviceExport]
[Differentiable]
float[RawPPISPRegLossIndexRQS::length] compute_raw_ppisp_rqs_regularization_loss(
float[PPISP_RQS_NUM_PARAMS] params
float[PPISP_RQS_NUM_PARAMS] params, bool exposure_arithmetic_mean
) {
PPISPParamsRQS p = get_ppisp_rqs_params(params);
float losses[RawPPISPRegLossIndexRQS::length] = { 0.0f };
// Exposure
losses[RawPPISPRegLossIndexRQS::SumExposure] = p.exposure;
losses[RawPPISPRegLossIndexRQS::SumExposure] =
ppisp_raw_exposure(p.exposure, exposure_arithmetic_mean);
// Vignetting
losses[RawPPISPRegLossIndexRQS::SumVignettingCrSquared] =
@@ -863,24 +875,24 @@ float[RawPPISPRegLossIndexRQS::length] compute_raw_ppisp_rqs_regularization_loss
[CudaDeviceExport]
float[PPISP_NUM_PARAMS] compute_raw_ppisp_regularization_loss_vjp(
float[PPISP_NUM_PARAMS] params,
float[PPISP_NUM_PARAMS] params, bool exposure_arithmetic_mean,
float[RawPPISPRegLossIndex::length] grad_out
) {
DifferentialPair<float[PPISP_NUM_PARAMS]> dp_params = diffPair(params);
bwd_diff(compute_raw_ppisp_regularization_loss)(
dp_params, grad_out
dp_params, exposure_arithmetic_mean, grad_out
);
return dp_params.d;
}
[CudaDeviceExport]
float[PPISP_RQS_NUM_PARAMS] compute_raw_ppisp_rqs_regularization_loss_vjp(
float[PPISP_RQS_NUM_PARAMS] params,
float[PPISP_RQS_NUM_PARAMS] params, bool exposure_arithmetic_mean,
float[RawPPISPRegLossIndexRQS::length] grad_out
) {
DifferentialPair<float[PPISP_RQS_NUM_PARAMS]> dp_params = diffPair(params);
bwd_diff(compute_raw_ppisp_rqs_regularization_loss)(
dp_params, grad_out
dp_params, exposure_arithmetic_mean, grad_out
);
return dp_params.d;
}
@@ -889,13 +901,14 @@ float[PPISP_RQS_NUM_PARAMS] compute_raw_ppisp_rqs_regularization_loss_vjp(
[CudaDeviceExport]
[Differentiable]
float[RawPPISPRegLossIndexNoCRF::length] compute_raw_ppisp_no_crf_regularization_loss(
float[PPISP_NO_CRF_NUM_PARAMS] params
float[PPISP_NO_CRF_NUM_PARAMS] params, bool exposure_arithmetic_mean
) {
PPISPParamsNoCRF p = get_ppisp_no_crf_params(params);
float losses[RawPPISPRegLossIndexNoCRF::length] = { 0.0f };
// Exposure
losses[RawPPISPRegLossIndexNoCRF::SumExposure] = p.exposure;
losses[RawPPISPRegLossIndexNoCRF::SumExposure] =
ppisp_raw_exposure(p.exposure, exposure_arithmetic_mean);
// Vignetting
losses[RawPPISPRegLossIndexNoCRF::SumVignettingCrSquared] =
@@ -948,12 +961,12 @@ float[RawPPISPRegLossIndexNoCRF::length] compute_raw_ppisp_no_crf_regularization
[CudaDeviceExport]
float[PPISP_NO_CRF_NUM_PARAMS] compute_raw_ppisp_no_crf_regularization_loss_vjp(
float[PPISP_NO_CRF_NUM_PARAMS] params,
float[PPISP_NO_CRF_NUM_PARAMS] params, bool exposure_arithmetic_mean,
float[RawPPISPRegLossIndexNoCRF::length] grad_out
) {
DifferentialPair<float[PPISP_NO_CRF_NUM_PARAMS]> dp_params = diffPair(params);
bwd_diff(compute_raw_ppisp_no_crf_regularization_loss)(
dp_params, grad_out
dp_params, exposure_arithmetic_mean, grad_out
);
return dp_params.d;
}
@@ -961,12 +974,13 @@ float[PPISP_NO_CRF_NUM_PARAMS] compute_raw_ppisp_no_crf_regularization_loss_vjp(
[CudaDeviceExport]
[Differentiable]
float[RawPPISPRegLossIndexNoCRFNoVig::length] compute_raw_ppisp_no_crf_no_vig_regularization_loss(
float[PPISP_NO_CRF_NO_VIG_NUM_PARAMS] params
float[PPISP_NO_CRF_NO_VIG_NUM_PARAMS] params, bool exposure_arithmetic_mean
) {
PPISPParamsNoCRFNoVig p = get_ppisp_no_crf_no_vig_params(params);
float losses[RawPPISPRegLossIndexNoCRFNoVig::length] = { 0.0f };
losses[RawPPISPRegLossIndexNoCRFNoVig::SumExposure] = p.exposure;
losses[RawPPISPRegLossIndexNoCRFNoVig::SumExposure] =
ppisp_raw_exposure(p.exposure, exposure_arithmetic_mean);
const static float2x2 zca_b = float2x2(0.0480542f, -0.0043631f, -0.0043631f, 0.0481283f);
const static float2x2 zca_r = float2x2(0.0580570f, -0.0179872f, -0.0179872f, 0.0431061f);
@@ -990,12 +1004,12 @@ float[RawPPISPRegLossIndexNoCRFNoVig::length] compute_raw_ppisp_no_crf_no_vig_re
[CudaDeviceExport]
float[PPISP_NO_CRF_NO_VIG_NUM_PARAMS] compute_raw_ppisp_no_crf_no_vig_regularization_loss_vjp(
float[PPISP_NO_CRF_NO_VIG_NUM_PARAMS] params,
float[PPISP_NO_CRF_NO_VIG_NUM_PARAMS] params, bool exposure_arithmetic_mean,
float[RawPPISPRegLossIndexNoCRFNoVig::length] grad_out
) {
DifferentialPair<float[PPISP_NO_CRF_NO_VIG_NUM_PARAMS]> dp_params = diffPair(params);
bwd_diff(compute_raw_ppisp_no_crf_no_vig_regularization_loss)(
dp_params, grad_out
dp_params, exposure_arithmetic_mean, grad_out
);
return dp_params.d;
}
@@ -1020,18 +1034,30 @@ float smooth_l1_loss(float x, float beta) {
}
}
// sum_exposure is ppisp_raw_exposure summed over the cameras.
[ForceInline]
[Differentiable]
float ppisp_exposure_mean_loss(float sum_exposure, int num_cameras,
bool arithmetic_mean) {
float mean = sum_exposure / num_cameras;
if (arithmetic_mean)
mean = log2(mean);
return smooth_l1_loss(mean, 0.1f);
}
[CudaDeviceExport]
[Differentiable]
float[PPISPRegLossIndex::length] compute_ppisp_regularization_loss(
float[RawPPISPRegLossIndex::length] raw_losses,
int num_cameras,
bool exposure_arithmetic_mean,
no_diff float[PPISPRegLossIndex::length] loss_weights
) {
float losses[PPISPRegLossIndex::length];
// Exposure mean
losses[PPISPRegLossIndex::ExposureMean] =
smooth_l1_loss(raw_losses[RawPPISPRegLossIndex::SumExposure] / num_cameras, 0.1f);
losses[PPISPRegLossIndex::ExposureMean] = ppisp_exposure_mean_loss(
raw_losses[RawPPISPRegLossIndex::SumExposure], num_cameras, exposure_arithmetic_mean);
// Vignetting center (cx, cy close to 0)
losses[PPISPRegLossIndex::VignettingCenter] =
@@ -1095,13 +1121,14 @@ float[PPISPRegLossIndex::length] compute_ppisp_regularization_loss(
float[PPISPRegLossIndex::length] compute_ppisp_rqs_regularization_loss(
float[RawPPISPRegLossIndexRQS::length] raw_losses,
int num_cameras,
bool exposure_arithmetic_mean,
no_diff float[PPISPRegLossIndex::length] loss_weights
) {
float losses[PPISPRegLossIndex::length];
// Exposure mean
losses[PPISPRegLossIndex::ExposureMean] =
smooth_l1_loss(raw_losses[RawPPISPRegLossIndexRQS::SumExposure] / num_cameras, 0.1f);
losses[PPISPRegLossIndex::ExposureMean] = ppisp_exposure_mean_loss(
raw_losses[RawPPISPRegLossIndexRQS::SumExposure], num_cameras, exposure_arithmetic_mean);
// Vignetting center (cx, cy close to 0)
losses[PPISPRegLossIndex::VignettingCenter] =
@@ -1164,12 +1191,14 @@ float[PPISPRegLossIndex::length] compute_ppisp_rqs_regularization_loss(
float[RawPPISPRegLossIndex::length] compute_ppisp_regularization_loss_vjp(
float[RawPPISPRegLossIndex::length] raw_losses,
int num_cameras,
bool exposure_arithmetic_mean,
float[PPISPRegLossIndex::length] loss_weights,
float[PPISPRegLossIndex::length] grad_out
) {
DifferentialPair<float[RawPPISPRegLossIndex::length]> dp_raw_losses = diffPair(raw_losses);
bwd_diff(compute_ppisp_regularization_loss)(
dp_raw_losses, num_cameras, loss_weights, grad_out
dp_raw_losses, num_cameras, exposure_arithmetic_mean, loss_weights,
grad_out
);
return dp_raw_losses.d;
}
@@ -1178,12 +1207,14 @@ float[RawPPISPRegLossIndex::length] compute_ppisp_regularization_loss_vjp(
float[RawPPISPRegLossIndexRQS::length] compute_ppisp_rqs_regularization_loss_vjp(
float[RawPPISPRegLossIndexRQS::length] raw_losses,
int num_cameras,
bool exposure_arithmetic_mean,
float[PPISPRegLossIndex::length] loss_weights,
float[PPISPRegLossIndex::length] grad_out
) {
DifferentialPair<float[RawPPISPRegLossIndexRQS::length]> dp_raw_losses = diffPair(raw_losses);
bwd_diff(compute_ppisp_rqs_regularization_loss)(
dp_raw_losses, num_cameras, loss_weights, grad_out
dp_raw_losses, num_cameras, exposure_arithmetic_mean, loss_weights,
grad_out
);
return dp_raw_losses.d;
}
@@ -1193,13 +1224,14 @@ float[RawPPISPRegLossIndexRQS::length] compute_ppisp_rqs_regularization_loss_vjp
float[PPISPRegLossIndex::length] compute_ppisp_no_crf_regularization_loss(
float[RawPPISPRegLossIndexNoCRF::length] raw_losses,
int num_cameras,
bool exposure_arithmetic_mean,
no_diff float[PPISPRegLossIndex::length] loss_weights
) {
float losses[PPISPRegLossIndex::length];
// Exposure mean
losses[PPISPRegLossIndex::ExposureMean] =
smooth_l1_loss(raw_losses[RawPPISPRegLossIndexNoCRF::SumExposure] / num_cameras, 0.1f);
losses[PPISPRegLossIndex::ExposureMean] = ppisp_exposure_mean_loss(
raw_losses[RawPPISPRegLossIndexNoCRF::SumExposure], num_cameras, exposure_arithmetic_mean);
// Vignetting center (cx, cy close to 0)
losses[PPISPRegLossIndex::VignettingCenter] =
@@ -1255,12 +1287,14 @@ float[PPISPRegLossIndex::length] compute_ppisp_no_crf_regularization_loss(
float[RawPPISPRegLossIndexNoCRF::length] compute_ppisp_no_crf_regularization_loss_vjp(
float[RawPPISPRegLossIndexNoCRF::length] raw_losses,
int num_cameras,
bool exposure_arithmetic_mean,
float[PPISPRegLossIndex::length] loss_weights,
float[PPISPRegLossIndex::length] grad_out
) {
DifferentialPair<float[RawPPISPRegLossIndexNoCRF::length]> dp_raw_losses = diffPair(raw_losses);
bwd_diff(compute_ppisp_no_crf_regularization_loss)(
dp_raw_losses, num_cameras, loss_weights, grad_out
dp_raw_losses, num_cameras, exposure_arithmetic_mean, loss_weights,
grad_out
);
return dp_raw_losses.d;
}
@@ -1270,13 +1304,14 @@ float[RawPPISPRegLossIndexNoCRF::length] compute_ppisp_no_crf_regularization_los
float[PPISPRegLossIndex::length] compute_ppisp_no_crf_no_vig_regularization_loss(
float[RawPPISPRegLossIndexNoCRFNoVig::length] raw_losses,
int num_cameras,
bool exposure_arithmetic_mean,
no_diff float[PPISPRegLossIndex::length] loss_weights
) {
float losses[PPISPRegLossIndex::length];
// Exposure mean
losses[PPISPRegLossIndex::ExposureMean] =
smooth_l1_loss(raw_losses[RawPPISPRegLossIndexNoCRFNoVig::SumExposure] / num_cameras, 0.1f);
losses[PPISPRegLossIndex::ExposureMean] = ppisp_exposure_mean_loss(
raw_losses[RawPPISPRegLossIndexNoCRFNoVig::SumExposure], num_cameras, exposure_arithmetic_mean);
// Color mean (latent offsets close to 0)
losses[PPISPRegLossIndex::ColorMean] = (
@@ -1309,12 +1344,14 @@ float[PPISPRegLossIndex::length] compute_ppisp_no_crf_no_vig_regularization_loss
float[RawPPISPRegLossIndexNoCRFNoVig::length] compute_ppisp_no_crf_no_vig_regularization_loss_vjp(
float[RawPPISPRegLossIndexNoCRFNoVig::length] raw_losses,
int num_cameras,
bool exposure_arithmetic_mean,
float[PPISPRegLossIndex::length] loss_weights,
float[PPISPRegLossIndex::length] grad_out
) {
DifferentialPair<float[RawPPISPRegLossIndexNoCRFNoVig::length]> dp_raw_losses = diffPair(raw_losses);
bwd_diff(compute_ppisp_no_crf_no_vig_regularization_loss)(
dp_raw_losses, num_cameras, loss_weights, grad_out
dp_raw_losses, num_cameras, exposure_arithmetic_mean, loss_weights,
grad_out
);
return dp_raw_losses.d;
}