MCPcopy Create free account
hub / github.com/MrNeRF/LichtFeld-Studio / initializeTrainingModel

Function initializeTrainingModel

src/training/training_setup.cpp:559–728  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

557 }
558
559 std::expected<void, std::string> initializeTrainingModel(
560 const lfs::core::param::TrainingParameters& params,
561 lfs::core::Scene& scene,
562 lfs::core::SplatTensorAllocator tensor_allocator) {
563
564 if (auto* model = scene.getTrainingModel()) {
565 applyTrainingSHDegree(*model, params.optimization.sh_degree);
566 if (auto result = appendAddedSplats(params, *model); !result) {
567 return result;
568 }
569 if (auto result = migrateTrainingModelToAllocator(params, *model, tensor_allocator); !result) {
570 return result;
571 }
572 scene.syncTrainingModelTopology(static_cast<size_t>(model->size()));
573 scene.notifyMutation(lfs::core::Scene::MutationType::MODEL_CHANGED);
574 return {};
575 }
576
577 lfs::core::NodeId point_cloud_node_id = lfs::core::NULL_NODE;
578 lfs::core::NodeId parent_id = lfs::core::NULL_NODE;
579 const lfs::core::PointCloud* point_cloud = nullptr;
580 glm::mat4 node_transform{1.0f};
581
582 for (const auto* node : scene.getNodes()) {
583 if (node->type == lfs::core::NodeType::POINTCLOUD && node->point_cloud) {
584 point_cloud_node_id = node->id;
585 parent_id = node->parent_id;
586 node_transform = node->transform();
587 point_cloud = node->point_cloud.get();
588 break;
589 }
590 }
591
592 lfs::core::PointCloud point_cloud_to_use;
593 const int max_cap = params.optimization.max_cap;
594
595 if (point_cloud && point_cloud->size() > 0) {
596 const lfs::core::CropBoxData* cropbox_data = nullptr;
597 lfs::core::NodeId cropbox_id = lfs::core::NULL_NODE;
598
599 if (point_cloud_node_id != lfs::core::NULL_NODE) {
600 cropbox_id = scene.getCropBoxForSplat(point_cloud_node_id);
601 if (cropbox_id != lfs::core::NULL_NODE) {
602 cropbox_data = scene.getCropBoxData(cropbox_id);
603 }
604 }
605
606 if (cropbox_data && cropbox_data->enabled) {
607 const glm::mat4 world_to_cropbox = glm::inverse(scene.getWorldTransform(cropbox_id));
608 const auto& means = point_cloud->means;
609 const auto& colors = point_cloud->colors;
610 const size_t num_points = point_cloud->size();
611
612 auto means_cpu = means.cpu();
613 auto colors_cpu = colors.cpu();
614 const float* means_ptr = means_cpu.ptr<float>();
615 const uint8_t* colors_ptr = colors_cpu.ptr<uint8_t>();
616

Callers 6

runHeadlessWithTCPFunction · 0.85
runHeadlessFunction · 0.85
load_datasetMethod · 0.85
startTrainingMethod · 0.85
TEST_FFunction · 0.85
TEST_PFunction · 0.85

Calls 15

applyTrainingSHDegreeFunction · 0.85
appendAddedSplatsFunction · 0.85
from_vectorFunction · 0.85
createRandomPointCloudFunction · 0.85
randomChoosePointCloudFunction · 0.85
formatFunction · 0.85
random_chooseFunction · 0.85
moveFunction · 0.85
getTrainingModelMethod · 0.80

Tested by 2

TEST_FFunction · 0.68
TEST_PFunction · 0.68