mirror of
https://github.com/harry7557558/spirula-studio.git
synced 2026-10-02 02:44:54 +08:00
use arithmetic mean for PPISP exposure centering
This commit is contained in:
+19
-11
@@ -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();
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
|
||||
|
||||
@@ -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];
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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.
|
||||
|
||||
+7781
-7250
File diff suppressed because it is too large
Load Diff
@@ -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("先做相機校正"),
|
||||
|
||||
@@ -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
@@ -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
@@ -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;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user