| 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 | |