-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdata.cpp
More file actions
75 lines (63 loc) · 2.92 KB
/
Copy pathdata.cpp
File metadata and controls
75 lines (63 loc) · 2.92 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
#include "data.h"
#include "nlohmann/json.hpp"
#define STB_IMAGE_IMPLEMENTATION
#include "stb_image.h"
#define STB_IMAGE_WRITE_IMPLEMENTATION
#include "stb_image_write.h"
#include <cmath>
#include <fstream>
#include <iostream>
#include <stdexcept>
#include <vector>
torch::Tensor load_image(const std::filesystem::path &path) {
int width = 0, height = 0;
uint8_t *data = stbi_load(path.string().c_str(), &width, &height, nullptr, 3);
if (!data) throw std::runtime_error("Could not read image '" + path.string() + "'");
auto image = torch::from_blob(data, {height, width, 3}, torch::kUInt8).clone();
stbi_image_free(data);
return image.to(torch::kFloat32) / 255.0f;
}
void save_image(const torch::Tensor &image, const std::filesystem::path &path) {
auto bytes = image.mul(255).clamp(0, 255).to(torch::kU8).to(torch::kCPU).contiguous();
const int height = bytes.size(0), width = bytes.size(1);
// Warn rather than throw: a failed preview must not abort a long training run.
if (stbi_write_png(path.string().c_str(), width, height, 3, bytes.data_ptr(),
width * 3) == 0) {
std::cerr << "Failed to save: " << path << std::endl;
}
}
Dataset load_dataset(const std::filesystem::path &json_path, int target_width) {
std::ifstream file(json_path);
if (!file.is_open()) {
throw std::runtime_error("Could not open '" + json_path.string() +
"'. The data path must be a directory containing "
"transforms.json and the image files it references.");
}
nlohmann::json data;
file >> data;
const auto dir = json_path.parent_path();
std::vector<torch::Tensor> images, poses;
for (const auto &frame : data["frames"]) {
images.push_back(load_image(dir / (frame["file_path"].get<std::string>() + ".png")));
std::array<float, 16> m{};
for (int i = 0; i < 4; i++)
for (int j = 0; j < 4; j++) m[i * 4 + j] = frame["transform_matrix"][i][j];
poses.push_back(torch::from_blob(m.data(), {4, 4}, torch::kFloat32).clone());
}
if (images.empty()) throw std::runtime_error("No frames in '" + json_path.string() + "'");
auto stacked = torch::stack(images);
const float aspect = static_cast<float>(stacked.size(2)) / stacked.size(1);
const int target_height = static_cast<int>(target_width * aspect);
if (stacked.size(1) != target_height || stacked.size(2) != target_width) {
stacked = torch::nn::functional::interpolate(
stacked.permute({0, 3, 1, 2}), // NHWC -> NCHW
torch::nn::functional::InterpolateFuncOptions()
.size(std::vector<int64_t>{target_height, target_width})
.mode(torch::kBilinear)
.align_corners(false))
.permute({0, 2, 3, 1});
}
const float camera_angle_x = data["camera_angle_x"];
return {stacked, torch::stack(poses),
0.5f * static_cast<float>(target_width) / std::tan(0.5f * camera_angle_x)};
}