-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy patheval.cpp
More file actions
77 lines (66 loc) · 2.87 KB
/
Copy patheval.cpp
File metadata and controls
77 lines (66 loc) · 2.87 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
#include "eval.h"
#include <cmath>
#include <fstream>
#include <iostream>
#include <limits>
namespace {
// Separable Gaussian as a depthwise conv kernel: [channels, 1, size, size].
torch::Tensor gaussian_kernel(int size, float sigma, int64_t channels,
const torch::TensorOptions &options) {
auto x = torch::arange(size, options) - (size - 1) / 2.0;
auto k = torch::exp(-x * x / (2.0f * sigma * sigma));
k = k / k.sum();
return torch::outer(k, k).expand({channels, 1, size, size}).contiguous();
}
// Single-scale SSIM with an 11x11 Gaussian window, averaged over channels.
float compute_ssim(const torch::Tensor &rendered, const torch::Tensor &target) {
constexpr int kWindow = 11;
constexpr float kC1 = 0.01f * 0.01f, kC2 = 0.03f * 0.03f; // (K * L)^2 with L = 1
const int64_t channels = rendered.size(2);
auto kernel = gaussian_kernel(kWindow, 1.5f, channels, rendered.options());
auto conv_opts =
torch::nn::functional::Conv2dFuncOptions().padding(kWindow / 2).groups(channels);
auto blur = [&](const torch::Tensor &t) {
return torch::nn::functional::conv2d(t, kernel, conv_opts);
};
auto x = rendered.permute({2, 0, 1}).unsqueeze(0).contiguous();
auto y = target.permute({2, 0, 1}).unsqueeze(0).contiguous();
auto mx = blur(x), my = blur(y);
auto mx2 = mx * mx, my2 = my * my, mxy = mx * my;
auto vx = blur(x * x) - mx2, vy = blur(y * y) - my2, vxy = blur(x * y) - mxy;
auto ssim = ((2.0f * mxy + kC1) * (2.0f * vxy + kC2)) /
((mx2 + my2 + kC1) * (vx + vy + kC2));
return ssim.mean().item<float>();
}
} // namespace
EvalMetrics compute_metrics(const torch::Tensor &rendered, const torch::Tensor &target) {
const float mse = torch::mse_loss(rendered, target).item<float>();
return {mse > 0.0f ? 10.0f * std::log10(1.0f / mse)
: std::numeric_limits<float>::infinity(),
std::sqrt(mse), compute_ssim(rendered, target)};
}
EvalMetrics mean_metrics(const std::vector<EvalMetrics> &views) {
EvalMetrics sum{0.0f, 0.0f, 0.0f};
for (const auto &m : views) {
sum.psnr += m.psnr;
sum.rmse += m.rmse;
sum.ssim += m.ssim;
}
const auto n = static_cast<float>(views.size());
return {sum.psnr / n, sum.rmse / n, sum.ssim / n};
}
void write_metrics_csv(const std::filesystem::path &path, int iter,
const std::vector<EvalMetrics> &views) {
const bool write_header = !std::filesystem::exists(path);
std::ofstream file(path, std::ios::app);
// Warn rather than throw: losing a metrics row must not abort a long training run.
if (!file.is_open()) {
std::cerr << "Failed to open metrics CSV: " << path << std::endl;
return;
}
if (write_header) file << "iter,view,psnr,rmse,ssim\n";
for (size_t v = 0; v < views.size(); v++) {
file << iter << "," << v << "," << views[v].psnr << "," << views[v].rmse << ","
<< views[v].ssim << "\n";
}
}