-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathutils.cpp
More file actions
98 lines (86 loc) · 4.02 KB
/
Copy pathutils.cpp
File metadata and controls
98 lines (86 loc) · 4.02 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
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
#include "utils.h"
#include "data.h"
#include <cmath>
#include <iostream>
#include <stdexcept>
void set_seed(int seed) {
torch::manual_seed(seed);
if (torch::cuda::is_available()) torch::cuda::manual_seed_all(seed);
}
torch::Device get_device(DevicePreference preference) {
if (preference == DevicePreference::CPU) {
std::cout << "Using CPU device" << std::endl;
return torch::kCPU;
}
if (torch::cuda::is_available()) {
std::cout << "Using CUDA device" << std::endl;
return torch::kCUDA;
}
if (preference == DevicePreference::CUDA) {
throw std::runtime_error(
"CUDA GPU is required for training, but LibTorch cannot see one. Check that an "
"NVIDIA driver is installed and running, that nvidia-smi works, and that this "
"binary is linked against a CUDA-enabled LibTorch build.");
}
std::cout << "CUDA unavailable; falling back to CPU device" << std::endl;
return torch::kCPU;
}
torch::Tensor spherical_pose(float azimuth, float elevation, float radius) {
const float phi = elevation * (M_PI / 180.0f);
const float theta = azimuth * (M_PI / 180.0f);
const float cp = std::cos(phi), sp = std::sin(phi);
const float ct = std::cos(theta), st = std::sin(theta);
auto translate = torch::tensor({{1.0f, 0.0f, 0.0f, 0.0f},
{0.0f, 1.0f, 0.0f, 0.0f},
{0.0f, 0.0f, 1.0f, radius},
{0.0f, 0.0f, 0.0f, 1.0f}});
auto rotate_phi = torch::tensor({{1.0f, 0.0f, 0.0f, 0.0f},
{0.0f, cp, -sp, 0.0f},
{0.0f, sp, cp, 0.0f},
{0.0f, 0.0f, 0.0f, 1.0f}});
auto rotate_theta = torch::tensor({{ct, 0.0f, -st, 0.0f},
{0.0f, 1.0f, 0.0f, 0.0f},
{st, 0.0f, ct, 0.0f},
{0.0f, 0.0f, 0.0f, 1.0f}});
// Axis flip into the Blender convention: camera looks down -z with +y up.
auto flip = torch::tensor({{-1.0f, 0.0f, 0.0f, 0.0f},
{0.0f, 0.0f, 1.0f, 0.0f},
{0.0f, 1.0f, 0.0f, 0.0f},
{0.0f, 0.0f, 0.0f, 1.0f}});
return flip.matmul(rotate_theta.matmul(rotate_phi.matmul(translate)));
}
void render_views(const NeRFRenderer &renderer, const std::string &prefix, int H, int W,
int n_frames, const std::filesystem::path &out_dir, float radius,
const RenderOptions &opt) {
std::cout << "Saving " << n_frames << " sample views..." << std::endl;
for (int i = 0; i < n_frames; i++) {
const auto azimuth = static_cast<float>(i) * 360.0f / static_cast<float>(n_frames);
auto pose = spherical_pose(azimuth, -30.0f, radius).to(renderer.device());
auto out = renderer.render_image(H, W, pose, opt);
const std::string suffix = prefix + "_" + std::to_string(i) + ".png";
save_image(out.rgb, out_dir / ("frame_" + suffix));
// Fixed-range normalisation keeps the depth scale stable across frames; near
// surfaces come out bright.
auto depth = 1.0 - ((out.depth - opt.z_near) / (opt.z_far - opt.z_near)).clamp(0.0, 1.0);
save_image(depth.unsqueeze(-1).expand({H, W, 3}), out_dir / ("frame_depth_" + suffix));
}
}
void save_checkpoint(const std::filesystem::path &path, const torch::nn::Module &model,
int iter) {
torch::serialize::OutputArchive archive;
model.save(archive);
archive.write("epoch", iter);
archive.save_to(path.string());
std::cout << "Model weights saved" << std::endl;
}
int load_checkpoint(const std::filesystem::path &path, torch::nn::Module &model) {
if (!std::filesystem::exists(path))
throw std::runtime_error("No checkpoint at '" + path.string() + "'");
torch::serialize::InputArchive archive;
archive.load_from(path.string());
model.load(archive);
c10::IValue iter;
archive.read("epoch", iter);
std::cout << "Loaded " << path << " from iteration " << iter.toInt() << std::endl;
return static_cast<int>(iter.toInt());
}