gsplat 1.0.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (116) hide show
  1. checksums.yaml +7 -0
  2. data/LICENSE.txt +202 -0
  3. data/README.md +236 -0
  4. data/docs/ACCEPTANCE.md +60 -0
  5. data/docs/BENCHMARKS.md +79 -0
  6. data/docs/DECISIONS.md +52 -0
  7. data/docs/MIGRATION.md +117 -0
  8. data/docs/PROFILE.md +75 -0
  9. data/docs/PROGRESS.md +199 -0
  10. data/docs/decisions/0000-template.md +17 -0
  11. data/docs/decisions/0001-pin-golden-reference-and-cpu-validation.md +27 -0
  12. data/docs/decisions/0002-retain-ruby-fallbacks-for-native-backend.md +26 -0
  13. data/docs/decisions/0003-share-projection-for-distorted-cameras.md +26 -0
  14. data/docs/decisions/0004-use-portable-world-space-reference-paths.md +26 -0
  15. data/docs/decisions/0005-reuse-ewa-core-for-2dgs.md +24 -0
  16. data/docs/decisions/0006-keep-eval3d-as-portable-reference.md +25 -0
  17. data/docs/decisions/0007-share-compositor-semantics-for-contribution-indices.md +25 -0
  18. data/examples/data/README.md +14 -0
  19. data/examples/data/colmap/images/view_000.png +0 -0
  20. data/examples/data/colmap/images/view_001.png +0 -0
  21. data/examples/data/colmap/images/view_002.png +0 -0
  22. data/examples/data/colmap/sparse/0/cameras.txt +2 -0
  23. data/examples/data/colmap/sparse/0/images.txt +6 -0
  24. data/examples/data/colmap/sparse/0/points3D.txt +17 -0
  25. data/examples/data/splats.ply +0 -0
  26. data/examples/fit_image.rb +35 -0
  27. data/examples/generate_sample_data.rb +147 -0
  28. data/examples/render_path.rb +117 -0
  29. data/examples/simple_trainer.rb +76 -0
  30. data/ext/gsplat_native/common.h +54 -0
  31. data/ext/gsplat_native/extconf.rb +39 -0
  32. data/ext/gsplat_native/gsplat_native.c +60 -0
  33. data/ext/gsplat_native/intersections.c +211 -0
  34. data/ext/gsplat_native/projection.c +199 -0
  35. data/ext/gsplat_native/raster_backward.c +129 -0
  36. data/ext/gsplat_native/raster_backward_bridge.c +69 -0
  37. data/ext/gsplat_native/raster_forward.c +150 -0
  38. data/ext/gsplat_native/rasterization.h +51 -0
  39. data/ext/gsplat_native/spherical_harmonics.c +129 -0
  40. data/gsplat.gemspec +36 -0
  41. data/lib/gsplat/autograd/context.rb +60 -0
  42. data/lib/gsplat/autograd/function.rb +68 -0
  43. data/lib/gsplat/autograd/variable.rb +159 -0
  44. data/lib/gsplat/backend/ruby/accumulate.rb +139 -0
  45. data/lib/gsplat/backend/ruby/accumulate_backward.rb +40 -0
  46. data/lib/gsplat/backend/ruby/eval3d_rasterizer.rb +175 -0
  47. data/lib/gsplat/backend/ruby/isect_tiles.rb +198 -0
  48. data/lib/gsplat/backend/ruby/projection.rb +251 -0
  49. data/lib/gsplat/backend/ruby/projection_backward.rb +190 -0
  50. data/lib/gsplat/backend/ruby/projection_covariance_vjp.rb +72 -0
  51. data/lib/gsplat/backend/ruby/projection_input_vjp.rb +252 -0
  52. data/lib/gsplat/backend/ruby/quat_scale_to_covar_preci.rb +139 -0
  53. data/lib/gsplat/backend/ruby/rasterize_to_indices_in_range.rb +121 -0
  54. data/lib/gsplat/backend/ruby/rasterize_to_pixels.rb +199 -0
  55. data/lib/gsplat/backend/ruby/rasterize_to_pixels_backward.rb +121 -0
  56. data/lib/gsplat/backend/ruby/spherical_harmonics.rb +135 -0
  57. data/lib/gsplat/backend/ruby/tile_compositor.rb +50 -0
  58. data/lib/gsplat/backend/ruby/tile_compositor_backward.rb +84 -0
  59. data/lib/gsplat/backend.rb +79 -0
  60. data/lib/gsplat/compression/grid_sort.rb +79 -0
  61. data/lib/gsplat/compression/kmeans.rb +121 -0
  62. data/lib/gsplat/compression/png.rb +147 -0
  63. data/lib/gsplat/compression/png_codec.rb +133 -0
  64. data/lib/gsplat/compression/quantizer.rb +79 -0
  65. data/lib/gsplat/io/checkpoint.rb +145 -0
  66. data/lib/gsplat/io/colmap.rb +175 -0
  67. data/lib/gsplat/io/colmap_binary.rb +98 -0
  68. data/lib/gsplat/io/colmap_text.rb +84 -0
  69. data/lib/gsplat/io/image.rb +63 -0
  70. data/lib/gsplat/io/image_backends.rb +81 -0
  71. data/lib/gsplat/io/npy.rb +189 -0
  72. data/lib/gsplat/io/ply.rb +185 -0
  73. data/lib/gsplat/io/ply_reader.rb +142 -0
  74. data/lib/gsplat/io/zip_archive.rb +183 -0
  75. data/lib/gsplat/math/camera_distortion.rb +123 -0
  76. data/lib/gsplat/math/camera_projection.rb +202 -0
  77. data/lib/gsplat/math/mat.rb +114 -0
  78. data/lib/gsplat/math/quaternion.rb +175 -0
  79. data/lib/gsplat/math/small_matrix_primitives.rb +112 -0
  80. data/lib/gsplat/math/spherical_harmonic_basis.rb +148 -0
  81. data/lib/gsplat/math/ssim.rb +153 -0
  82. data/lib/gsplat/native.rb +30 -0
  83. data/lib/gsplat/native_ops.rb +148 -0
  84. data/lib/gsplat/native_raster_ops.rb +76 -0
  85. data/lib/gsplat/ops/accumulate.rb +68 -0
  86. data/lib/gsplat/ops/eval3d_rasterize.rb +62 -0
  87. data/lib/gsplat/ops/isect_tiles.rb +38 -0
  88. data/lib/gsplat/ops/projection.rb +154 -0
  89. data/lib/gsplat/ops/quat_scale_to_covar_preci.rb +118 -0
  90. data/lib/gsplat/ops/rasterize_to_indices_in_range.rb +38 -0
  91. data/lib/gsplat/ops/rasterize_to_pixels.rb +105 -0
  92. data/lib/gsplat/ops/relocation.rb +94 -0
  93. data/lib/gsplat/ops/spherical_harmonics.rb +71 -0
  94. data/lib/gsplat/ops/tensor_shape_ops.rb +173 -0
  95. data/lib/gsplat/ops/tensor_value_ops.rb +150 -0
  96. data/lib/gsplat/optim/adam.rb +212 -0
  97. data/lib/gsplat/optim/lr_scheduler.rb +36 -0
  98. data/lib/gsplat/optim/selective_adam.rb +68 -0
  99. data/lib/gsplat/rasterization.rb +130 -0
  100. data/lib/gsplat/rasterization_2dgs.rb +140 -0
  101. data/lib/gsplat/rasterization_helpers.rb +162 -0
  102. data/lib/gsplat/rasterization_validation.rb +99 -0
  103. data/lib/gsplat/strategy/base.rb +49 -0
  104. data/lib/gsplat/strategy/default.rb +188 -0
  105. data/lib/gsplat/strategy/mcmc.rb +103 -0
  106. data/lib/gsplat/strategy/mcmc_ops.rb +143 -0
  107. data/lib/gsplat/strategy/ops.rb +165 -0
  108. data/lib/gsplat/training/config.rb +91 -0
  109. data/lib/gsplat/training/image_fitter.rb +158 -0
  110. data/lib/gsplat/training/losses.rb +191 -0
  111. data/lib/gsplat/training/scene.rb +116 -0
  112. data/lib/gsplat/training/trainer.rb +236 -0
  113. data/lib/gsplat/utils.rb +110 -0
  114. data/lib/gsplat/version.rb +6 -0
  115. data/lib/gsplat.rb +98 -0
  116. metadata +181 -0
@@ -0,0 +1,188 @@
1
+ # frozen_string_literal: true
2
+
3
+ module Gsplat
4
+ module Strategy
5
+ # Original 3DGS densification strategy with gsplat refinements.
6
+ class Default < Base
7
+ # Upstream-compatible densification defaults.
8
+ DEFAULTS = {
9
+ prune_opa: 0.005,
10
+ grow_grad2d: 0.0002,
11
+ grow_scale3d: 0.01,
12
+ grow_scale2d: 0.05,
13
+ prune_scale3d: 0.1,
14
+ prune_scale2d: 0.15,
15
+ refine_scale2d_stop_iter: 0,
16
+ refine_start_iter: 500,
17
+ refine_stop_iter: 15_000,
18
+ reset_every: 3_000,
19
+ refine_every: 100,
20
+ pause_refine_after_reset: 0,
21
+ absgrad: false,
22
+ revised_opacity: false,
23
+ key_for_gradient: :means2d
24
+ }.freeze
25
+
26
+ attr_reader :absgrad, :grow_grad2d, :grow_scale2d, :grow_scale3d,
27
+ :key_for_gradient, :pause_refine_after_reset, :prune_opa,
28
+ :prune_scale2d, :prune_scale3d, :refine_every,
29
+ :refine_scale2d_stop_iter, :refine_start_iter,
30
+ :refine_stop_iter, :reset_every, :revised_opacity
31
+
32
+ def initialize(**options)
33
+ super()
34
+ unknown = options.keys - DEFAULTS.keys
35
+ raise ArgumentError, "unknown options: #{unknown.join(', ')}" unless unknown.empty?
36
+
37
+ DEFAULTS.merge(options).each { |name, value| instance_variable_set(:"@#{name}", value) }
38
+ validate_options!
39
+ end
40
+
41
+ # Creates mutable gradient/radius statistics for a scene.
42
+ #
43
+ # @param scene_scale [Numeric]
44
+ # @return [Hash]
45
+ def initialize_state(scene_scale:)
46
+ super.merge(grad2d: nil, count: nil, radii: nil)
47
+ end
48
+
49
+ # rubocop:disable Metrics/ParameterLists
50
+ def step_post_backward(params:, optimizers:, state:, step:, info:, **_options)
51
+ check_sanity(params, optimizers)
52
+ collect_statistics!(params, state, info)
53
+ refine!(params, optimizers, state, step) if should_refine?(step)
54
+ reset_due = step.positive? && (step % reset_every).zero?
55
+ Ops.reset_opacity!(params, optimizers, maximum: 2 * prune_opa) if reset_due
56
+ state
57
+ end
58
+ # rubocop:enable Metrics/ParameterLists
59
+
60
+ # rubocop:disable Metrics/AbcSize
61
+ def refinement_masks(params, state)
62
+ count = params.fetch(:means).data.shape[0]
63
+ average = params.fetch(:means).data.class.zeros(count)
64
+ observed = state[:count].gt(0)
65
+ average[observed] = state[:grad2d][observed] / state[:count][observed] if observed.any?
66
+ high = average.gt(grow_grad2d)
67
+ max_scale = Numo::NMath.exp(params.fetch(:scales).data).max(axis: 1)
68
+ small = max_scale.le(grow_scale3d * state.fetch(:scene_scale))
69
+ small &= state[:radii].le(grow_scale2d) if refine_scale2d_stop_iter.positive? && state[:radii]
70
+ [high & small, high & small.eq(0)]
71
+ end
72
+ # rubocop:enable Metrics/AbcSize
73
+
74
+ private
75
+
76
+ def validate_options!
77
+ valid_opacity = prune_opa.positive? && prune_opa < 0.5
78
+ raise ArgumentError, "prune_opa must be in (0,0.5)" unless valid_opacity
79
+
80
+ intervals = [refine_every, reset_every]
81
+ valid_intervals = intervals.all? { |value| value.is_a?(Integer) && value.positive? }
82
+ raise ArgumentError, "refinement intervals must be positive" unless valid_intervals
83
+ return if refine_start_iter <= refine_stop_iter
84
+
85
+ raise ArgumentError, "refine_start_iter must not exceed refine_stop_iter"
86
+ end
87
+
88
+ # rubocop:disable Metrics/AbcSize
89
+ def collect_statistics!(params, state, info)
90
+ projected = info.fetch(key_for_gradient)
91
+ gradient = if absgrad
92
+ projected.respond_to?(:absgrad) ? projected.absgrad : info[:means2d_absgrad]
93
+ else
94
+ projected.respond_to?(:grad) ? projected.grad : nil
95
+ end
96
+ return unless gradient
97
+
98
+ count = params.fetch(:means).data.shape[0]
99
+ initialize_statistics!(state, count, gradient.class)
100
+ width = info.fetch(:width)
101
+ height = info.fetch(:height)
102
+ normalized = gradient.dup
103
+ normalized[true, true, 0] *= width / 2.0
104
+ normalized[true, true, 1] *= height / 2.0
105
+ norm = (normalized**2).sum(axis: 2)**0.5
106
+ radii = Gsplat::Ops::TensorOps.data(info.fetch(:radii))
107
+ radii = radii.max(axis: radii.ndim - 1) if radii.ndim == 3 && radii.shape[-1] == 2
108
+ visible = radii.gt(0)
109
+ norm[visible.eq(0)] = 0
110
+ state[:grad2d] += norm.sum(axis: 0)
111
+ state[:count] += gradient.class.cast(visible).sum(axis: 0)
112
+ normalized_radii = radii / [width, height].max.to_f
113
+ camera_max = normalized_radii.max(axis: 0)
114
+ larger = camera_max.gt(state[:radii])
115
+ state[:radii][larger] = camera_max[larger] if larger.any?
116
+ end
117
+ # rubocop:enable Metrics/AbcSize
118
+
119
+ def initialize_statistics!(state, count, type)
120
+ return if state[:grad2d]&.shape == [count]
121
+
122
+ state[:grad2d] = type.zeros(count)
123
+ state[:count] = type.zeros(count)
124
+ state[:radii] = type.zeros(count)
125
+ end
126
+
127
+ def should_refine?(step)
128
+ return false unless step.between?(refine_start_iter, refine_stop_iter)
129
+ return false unless (step % refine_every).zero?
130
+ return true if pause_refine_after_reset.zero?
131
+
132
+ (step % reset_every) > pause_refine_after_reset
133
+ end
134
+
135
+ def refine!(params, optimizers, state, step)
136
+ duplicate_mask, split_mask = refinement_masks(params, state)
137
+ duplicate_count = Ops.duplicate!(params, optimizers, duplicate_mask)
138
+ revise_duplicated_opacity!(params, duplicate_mask, duplicate_count) if revised_opacity
139
+ extended_split = extend_mask(split_mask, duplicate_count)
140
+ Ops.split!(params, optimizers, extended_split)
141
+ prune!(params, optimizers, state, step)
142
+ initialize_statistics!(
143
+ state,
144
+ params.fetch(:means).data.shape[0],
145
+ params.fetch(:means).data.class
146
+ )
147
+ end
148
+
149
+ def revise_duplicated_opacity!(params, original_mask, duplicate_count)
150
+ return if duplicate_count.zero?
151
+
152
+ variable = params.fetch(:opacities)
153
+ indices = original_mask.where.to_a
154
+ active = sigmoid(variable.data[indices])
155
+ revised = 1 - Numo::NMath.sqrt(1 - active)
156
+ logits = Numo::NMath.log(revised / (1 - revised))
157
+ variable.data[indices] = logits
158
+ variable.data[-duplicate_count..] = logits
159
+ end
160
+
161
+ def extend_mask(mask, count)
162
+ return mask if count.zero?
163
+
164
+ output = Numo::Bit.zeros(mask.size + count)
165
+ output[0...mask.size] = mask
166
+ output
167
+ end
168
+
169
+ def prune!(params, optimizers, state, step)
170
+ opacity = sigmoid(params.fetch(:opacities).data)
171
+ prune = opacity.lt(prune_opa)
172
+ if step > reset_every
173
+ max_scale = Numo::NMath.exp(params.fetch(:scales).data).max(axis: 1)
174
+ prune |= max_scale.gt(prune_scale3d * state.fetch(:scene_scale))
175
+ if refine_scale2d_stop_iter.positive? && step < refine_scale2d_stop_iter &&
176
+ state[:radii]&.shape == prune.shape
177
+ prune |= state[:radii].gt(prune_scale2d)
178
+ end
179
+ end
180
+ Ops.remove!(params, optimizers, prune)
181
+ end
182
+
183
+ def sigmoid(values)
184
+ 1.0 / (1 + Numo::NMath.exp(-values))
185
+ end
186
+ end
187
+ end
188
+ end
@@ -0,0 +1,103 @@
1
+ # frozen_string_literal: true
2
+
3
+ module Gsplat
4
+ module Strategy
5
+ # 3D Gaussian Splatting as Markov Chain Monte Carlo strategy.
6
+ class MCMC < Base
7
+ # Upstream-compatible relocation and noise defaults.
8
+ DEFAULTS = {
9
+ cap_max: 1_000_000,
10
+ noise_lr: 5e5,
11
+ refine_start_iter: 500,
12
+ refine_stop_iter: 25_000,
13
+ refine_every: 100,
14
+ min_opacity: 0.005,
15
+ verbose: false
16
+ }.freeze
17
+
18
+ attr_reader :cap_max, :min_opacity, :noise_lr, :refine_every,
19
+ :refine_start_iter, :refine_stop_iter, :verbose
20
+
21
+ def initialize(**options)
22
+ super()
23
+ unknown = options.keys - DEFAULTS.keys
24
+ raise ArgumentError, "unknown options: #{unknown.join(', ')}" unless unknown.empty?
25
+
26
+ DEFAULTS.merge(options).each { |name, value| instance_variable_set(:"@#{name}", value) }
27
+ validate_options!
28
+ end
29
+
30
+ # Creates strategy state including the relocation binomial table.
31
+ #
32
+ # @param scene_scale [Numeric]
33
+ # @return [Hash]
34
+ def initialize_state(scene_scale: 1.0)
35
+ super.merge(binoms: Gsplat::Ops::Relocation.binomial_table(n_max: 51))
36
+ end
37
+
38
+ # rubocop:disable Metrics/ParameterLists
39
+ def step_post_backward(params:, optimizers:, state:, step:, info:, **options)
40
+ raise ArgumentError, "info must be a Hash" unless info.is_a?(Hash)
41
+
42
+ learning_rate = options.fetch(:lr)
43
+ check_sanity(params, optimizers)
44
+ if should_refine?(step)
45
+ relocated = relocate_dead!(params, optimizers, state)
46
+ added = add_new!(params, optimizers, state)
47
+ log_step(step, relocated, added) if verbose
48
+ end
49
+ Ops.inject_position_noise!(params, scaler: learning_rate * noise_lr)
50
+ state
51
+ end
52
+ # rubocop:enable Metrics/ParameterLists
53
+
54
+ private
55
+
56
+ def validate_options!
57
+ raise ArgumentError, "cap_max must be positive" unless cap_max.is_a?(Integer) && cap_max.positive?
58
+ raise ArgumentError, "noise_lr must be non-negative" unless noise_lr >= 0
59
+ unless refine_every.is_a?(Integer) && refine_every.positive?
60
+ raise ArgumentError, "refine_every must be positive"
61
+ end
62
+ raise ArgumentError, "min_opacity must be in (0,1)" unless min_opacity.positive? && min_opacity < 1
63
+ return if refine_start_iter <= refine_stop_iter
64
+
65
+ raise ArgumentError, "refine_start_iter must not exceed refine_stop_iter"
66
+ end
67
+
68
+ def should_refine?(step)
69
+ step < refine_stop_iter && step > refine_start_iter && (step % refine_every).zero?
70
+ end
71
+
72
+ def relocate_dead!(params, optimizers, state)
73
+ opacity = 1.0 / (1.0 + Numo::NMath.exp(-params.fetch(:opacities).data))
74
+ dead = opacity.le(min_opacity)
75
+ Ops.relocate!(
76
+ params,
77
+ optimizers,
78
+ dead,
79
+ binoms: state.fetch(:binoms),
80
+ min_opacity: min_opacity
81
+ )
82
+ end
83
+
84
+ def add_new!(params, optimizers, state)
85
+ current = params.fetch(:means).data.shape[0]
86
+ target = [cap_max, (1.05 * current).to_i].min
87
+ Ops.sample_add!(
88
+ params,
89
+ optimizers,
90
+ [target - current, 0].max,
91
+ binoms: state.fetch(:binoms),
92
+ min_opacity: min_opacity
93
+ )
94
+ end
95
+
96
+ def log_step(step, relocated, added)
97
+ Gsplat.logger.info(
98
+ "MCMC step #{step}: relocated #{relocated}, added #{added}"
99
+ )
100
+ end
101
+ end
102
+ end
103
+ end
@@ -0,0 +1,143 @@
1
+ # frozen_string_literal: true
2
+
3
+ module Gsplat
4
+ module Strategy
5
+ # MCMC-specific parameter edits synchronized with Adam state.
6
+ module Ops
7
+ module_function
8
+
9
+ # rubocop:disable Metrics/AbcSize
10
+ def relocate!(params, optimizers, mask, binoms:, **options)
11
+ min_opacity = options.fetch(:min_opacity, 0.005)
12
+ rng = options.fetch(:rng, Gsplat.rng)
13
+ count = params.fetch(:means).data.shape[0]
14
+ validate_mcmc_mask!(mask, count)
15
+ dead = Numo::Bit.cast(mask).where.to_a
16
+ return 0 if dead.empty?
17
+
18
+ alive = Numo::Bit.cast(mask).eq(0).where.to_a
19
+ raise Gsplat::Error, "cannot relocate when every Gaussian is dead" if alive.empty?
20
+
21
+ weights = sigmoid_mcmc(params.fetch(:opacities).data[alive])
22
+ sampled = weighted_indices(weights, dead.size, rng).map { |index| alive.fetch(index) }
23
+ update_relocation_sources!(params, sampled, binoms, min_opacity)
24
+ prepare_mcmc_optimizers!(optimizers)
25
+ params.each do |key, variable|
26
+ variable.data[*([dead] + Array.new(variable.data.ndim - 1, true))] =
27
+ variable.data[*([sampled] + Array.new(variable.data.ndim - 1, true))]
28
+ optimizers.fetch(key).zero_state_at!(sampled)
29
+ optimizers.fetch(key).zero_state_at!(dead)
30
+ end
31
+ dead.size
32
+ end
33
+ # rubocop:enable Metrics/AbcSize
34
+
35
+ def sample_add!(params, optimizers, count, binoms:, **options)
36
+ min_opacity = options.fetch(:min_opacity, 0.005)
37
+ rng = options.fetch(:rng, Gsplat.rng)
38
+ raise ArgumentError, "count must be a non-negative integer" unless count.is_a?(Integer) && !count.negative?
39
+ return 0 if count.zero?
40
+
41
+ weights = sigmoid_mcmc(params.fetch(:opacities).data)
42
+ sampled = weighted_indices(weights, count, rng)
43
+ update_relocation_sources!(params, sampled, binoms, min_opacity)
44
+ prepare_mcmc_optimizers!(optimizers)
45
+ params.each do |key, variable|
46
+ rows = variable.data[*([sampled] + Array.new(variable.data.ndim - 1, true))].dup
47
+ variable.replace_data!(append_mcmc_rows(variable.data, rows))
48
+ optimizers.fetch(key).append!(count)
49
+ end
50
+ count
51
+ end
52
+
53
+ # rubocop:disable Metrics/AbcSize
54
+ def inject_position_noise!(params, scaler:, rng: Gsplat.rng)
55
+ return params if scaler.zero?
56
+
57
+ opacities = sigmoid_mcmc(params.fetch(:opacities).data)
58
+ scales = Numo::NMath.exp(params.fetch(:scales).data)
59
+ covariances, = Gsplat.quat_scale_to_covar_preci(
60
+ params.fetch(:quats).data,
61
+ scales,
62
+ compute_preci: false
63
+ )
64
+ noise = params.fetch(:means).data.class.zeros(*params.fetch(:means).data.shape)
65
+ noise.shape[0].times do |index|
66
+ gate = 1.0 / (1.0 + ::Math.exp(-100.0 * ((1.0 - opacities[index].to_f) - 0.995)))
67
+ standard = noise.class.cast(Array.new(3) { normal_mcmc_sample(rng) })
68
+ noise[index, true] = covariances[index, true, true].dot(standard) * gate * scaler
69
+ end
70
+ params.fetch(:means).data[] = params.fetch(:means).data + noise
71
+ params
72
+ end
73
+ # rubocop:enable Metrics/AbcSize
74
+
75
+ def update_relocation_sources!(params, sampled, binoms, min_opacity)
76
+ source_opacities = sigmoid_mcmc(params.fetch(:opacities).data[sampled])
77
+ source_scales = Numo::NMath.exp(params.fetch(:scales).data[sampled, true])
78
+ counts = sampled.tally
79
+ ratios = Numo::Int32.cast(sampled.map { |index| counts.fetch(index) + 1 })
80
+ opacities, scales = Gsplat.relocation(source_opacities, source_scales, ratios, binoms: binoms)
81
+ epsilon = Numo::SFloat::EPSILON
82
+ opacities = clamp_mcmc(opacities, min_opacity, 1.0 - epsilon)
83
+ params.fetch(:opacities).data[sampled] = Numo::NMath.log(opacities / (1.0 - opacities))
84
+ params.fetch(:scales).data[sampled, true] = Numo::NMath.log(scales)
85
+ end
86
+ private_class_method :update_relocation_sources!
87
+
88
+ def weighted_indices(weights, count, rng)
89
+ cumulative = weights.to_a.each_with_object([]) do |weight, sums|
90
+ sums << (weight.to_f + (sums.last || 0.0))
91
+ end
92
+ valid_sum = cumulative.last&.positive? && cumulative.last.finite?
93
+ raise Gsplat::Error, "sampling weights must have positive finite sum" unless valid_sum
94
+
95
+ Array.new(count) do
96
+ draw = rng.rand * cumulative.last
97
+ cumulative.bsearch_index { |value| value > draw } || (cumulative.length - 1)
98
+ end
99
+ end
100
+ private_class_method :weighted_indices
101
+
102
+ def append_mcmc_rows(array, rows)
103
+ output = array.class.zeros(*([array.shape[0] + rows.shape[0]] + array.shape[1..]))
104
+ output[*([0...array.shape[0]] + Array.new(array.ndim - 1, true))] = array
105
+ output[*([array.shape[0]...output.shape[0]] + Array.new(array.ndim - 1, true))] = rows
106
+ output
107
+ end
108
+ private_class_method :append_mcmc_rows
109
+
110
+ def prepare_mcmc_optimizers!(optimizers)
111
+ optimizers.each_value(&:state)
112
+ end
113
+ private_class_method :prepare_mcmc_optimizers!
114
+
115
+ def validate_mcmc_mask!(mask, count)
116
+ return if mask.is_a?(Numo::NArray) && mask.shape == [count]
117
+
118
+ actual = mask.respond_to?(:shape) ? mask.shape.inspect : mask.class.to_s
119
+ raise ShapeError, "expected mask [#{count}], got #{actual}"
120
+ end
121
+ private_class_method :validate_mcmc_mask!
122
+
123
+ def sigmoid_mcmc(values)
124
+ 1.0 / (1.0 + Numo::NMath.exp(-values))
125
+ end
126
+ private_class_method :sigmoid_mcmc
127
+
128
+ def clamp_mcmc(values, minimum, maximum)
129
+ output = values.dup
130
+ output[output.lt(minimum)] = minimum
131
+ output[output.gt(maximum)] = maximum
132
+ output
133
+ end
134
+ private_class_method :clamp_mcmc
135
+
136
+ def normal_mcmc_sample(rng)
137
+ radius = ::Math.sqrt(-2.0 * ::Math.log([rng.rand, Float::MIN].max))
138
+ radius * ::Math.cos(2.0 * ::Math::PI * rng.rand)
139
+ end
140
+ private_class_method :normal_mcmc_sample
141
+ end
142
+ end
143
+ end
@@ -0,0 +1,165 @@
1
+ # frozen_string_literal: true
2
+
3
+ module Gsplat
4
+ # Densification and relocation policies for training.
5
+ module Strategy
6
+ # Structural parameter edits synchronized with Adam moment buffers.
7
+ module Ops
8
+ module_function
9
+
10
+ # Appends exact copies of selected Gaussian rows and zeroed moments.
11
+ #
12
+ # @return [Integer] number of rows appended
13
+ def duplicate!(params, optimizers, mask)
14
+ indices = selected_indices(mask, parameter_count(params))
15
+ return 0 if indices.empty?
16
+
17
+ prepare_optimizers!(optimizers)
18
+ params.each do |key, variable|
19
+ appended = select_rows(variable.data, indices)
20
+ variable.replace_data!(append_rows(variable.data, appended))
21
+ optimizers.fetch(key).append!(indices.size)
22
+ end
23
+ indices.size
24
+ end
25
+
26
+ # Removes selected Gaussian rows and matching optimizer moments.
27
+ #
28
+ # @return [Integer] number of rows removed
29
+ def remove!(params, optimizers, mask)
30
+ validate_mask!(mask, parameter_count(params))
31
+ keep = Numo::Bit.cast(mask).eq(0)
32
+ removed = mask.count_true
33
+ return 0 if removed.zero?
34
+
35
+ prepare_optimizers!(optimizers)
36
+ params.each do |key, variable|
37
+ variable.replace_data!(select_rows(variable.data, keep))
38
+ optimizers.fetch(key).select!(keep)
39
+ end
40
+ removed
41
+ end
42
+
43
+ # Replaces selected Gaussians with two sampled children.
44
+ #
45
+ # @return [Integer] number of parents split
46
+ def split!(params, optimizers, mask, rng: Gsplat.rng)
47
+ count = parameter_count(params)
48
+ indices = selected_indices(mask, count)
49
+ return 0 if indices.empty?
50
+
51
+ children = split_children(params, indices, rng)
52
+ keep = Numo::Bit.cast(mask).eq(0)
53
+ prepare_optimizers!(optimizers)
54
+ params.each do |key, variable|
55
+ retained = select_rows(variable.data, keep)
56
+ variable.replace_data!(append_rows(retained, children.fetch(key)))
57
+ optimizer = optimizers.fetch(key)
58
+ optimizer.select!(keep)
59
+ optimizer.append!(indices.size * 2)
60
+ end
61
+ indices.size
62
+ end
63
+
64
+ def reset_opacity!(params, optimizers, maximum:)
65
+ raise ArgumentError, "maximum opacity must be in (0,1)" unless maximum.positive? && maximum < 1
66
+
67
+ variable = params.fetch(:opacities)
68
+ cap = ::Math.log(maximum / (1 - maximum))
69
+ values = variable.data.dup
70
+ above = values.gt(cap)
71
+ values[above] = cap if above.any?
72
+ variable.replace_data!(values)
73
+ optimizer = optimizers.fetch(:opacities)
74
+ optimizer.state
75
+ optimizer.zero_state_at!(0...values.shape[0])
76
+ above.count_true
77
+ end
78
+
79
+ def parameter_count(params)
80
+ raise ArgumentError, "params must not be empty" if params.empty?
81
+
82
+ params.values.first.data.shape[0]
83
+ end
84
+ private_class_method :parameter_count
85
+
86
+ def selected_indices(mask, count)
87
+ validate_mask!(mask, count)
88
+ Numo::Bit.cast(mask).where.to_a
89
+ end
90
+ private_class_method :selected_indices
91
+
92
+ def validate_mask!(mask, count)
93
+ return if mask.is_a?(Numo::NArray) && mask.shape == [count]
94
+
95
+ actual = mask.respond_to?(:shape) ? mask.shape.inspect : mask.class.to_s
96
+ raise ShapeError, "expected mask [#{count}], got #{actual}"
97
+ end
98
+ private_class_method :validate_mask!
99
+
100
+ def prepare_optimizers!(optimizers)
101
+ optimizers.each_value(&:state)
102
+ end
103
+ private_class_method :prepare_optimizers!
104
+
105
+ def select_rows(array, selection)
106
+ array[*([selection] + Array.new(array.ndim - 1, true))].dup
107
+ end
108
+ private_class_method :select_rows
109
+
110
+ def append_rows(array, rows)
111
+ output = array.class.zeros(*([array.shape[0] + rows.shape[0]] + array.shape[1..]))
112
+ output[*([0...array.shape[0]] + Array.new(array.ndim - 1, true))] = array
113
+ output[*([array.shape[0]...output.shape[0]] + Array.new(array.ndim - 1, true))] = rows
114
+ output
115
+ end
116
+ private_class_method :append_rows
117
+
118
+ def split_children(params, indices, rng)
119
+ children = params.to_h do |key, variable|
120
+ selected = select_rows(variable.data, indices)
121
+ doubled = repeat_rows(selected, 2)
122
+ [key, doubled]
123
+ end
124
+ children[:scales] -= ::Math.log(1.6)
125
+ children[:means] += split_offsets(params, indices, rng)
126
+ children
127
+ end
128
+ private_class_method :split_children
129
+
130
+ def repeat_rows(array, repeat)
131
+ output = array.class.zeros(*([array.shape[0] * repeat] + array.shape[1..]))
132
+ array.shape[0].times do |index|
133
+ repeat.times do |copy|
134
+ output[*([(index * repeat) + copy] + Array.new(array.ndim - 1, true))] =
135
+ array[*([index] + Array.new(array.ndim - 1, true))]
136
+ end
137
+ end
138
+ output
139
+ end
140
+ private_class_method :repeat_rows
141
+
142
+ def split_offsets(params, indices, rng)
143
+ quaternions = select_rows(params.fetch(:quats).data, indices)
144
+ log_scales = select_rows(params.fetch(:scales).data, indices)
145
+ rotations = Math::Quaternion.to_rotmat(quaternions)
146
+ output = params.fetch(:means).data.class.zeros(indices.size * 2, 3)
147
+ indices.size.times do |index|
148
+ 2.times do |child|
149
+ noise = output.class.cast(Array.new(3) { normal_sample(rng) })
150
+ scaled = noise * Numo::NMath.exp(log_scales[index, true])
151
+ output[(index * 2) + child, true] = rotations[index, true, true].dot(scaled)
152
+ end
153
+ end
154
+ output
155
+ end
156
+ private_class_method :split_offsets
157
+
158
+ def normal_sample(rng)
159
+ radius = ::Math.sqrt(-2 * ::Math.log([rng.rand, Float::MIN].max))
160
+ radius * ::Math.cos(2 * ::Math::PI * rng.rand)
161
+ end
162
+ private_class_method :normal_sample
163
+ end
164
+ end
165
+ end
@@ -0,0 +1,91 @@
1
+ # frozen_string_literal: true
2
+
3
+ module Gsplat
4
+ module Training
5
+ # Trainer hyperparameters matching the upstream simple trainer defaults.
6
+ class Config
7
+ # Default training, optimizer, evaluation, and output options.
8
+ DEFAULTS = {
9
+ max_steps: 30_000,
10
+ batch_size: 1,
11
+ sh_degree: 3,
12
+ sh_degree_interval: 1_000,
13
+ ssim_lambda: 0.2,
14
+ init_opacity: 0.1,
15
+ init_scale: 1.0,
16
+ means_lr: 1.6e-4,
17
+ means_lr_final: 1.6e-6,
18
+ scales_lr: 5e-3,
19
+ quats_lr: 1e-3,
20
+ opacities_lr: 5e-2,
21
+ sh0_lr: 2.5e-3,
22
+ shN_lr: 1.25e-4,
23
+ opacity_reg: 0.0,
24
+ scale_reg: 0.0,
25
+ eval_steps: [7_000, 30_000],
26
+ save_steps: [7_000, 30_000],
27
+ log_every: 100,
28
+ random_background: false,
29
+ near_plane: 0.01,
30
+ far_plane: 1e10,
31
+ tile_size: 16,
32
+ rasterize_mode: "classic",
33
+ model_type: :three_d,
34
+ output_dir: "results",
35
+ seed: 42
36
+ }.freeze
37
+
38
+ # rubocop:disable Naming/MethodName
39
+ attr_reader :batch_size, :eval_steps, :far_plane, :init_opacity,
40
+ :init_scale, :log_every, :max_steps, :means_lr,
41
+ :means_lr_final, :model_type, :near_plane, :opacities_lr,
42
+ :opacity_reg, :output_dir, :quats_lr, :random_background,
43
+ :rasterize_mode, :save_steps, :scales_lr, :scale_reg, :seed,
44
+ :sh0_lr, :shN_lr, :sh_degree, :sh_degree_interval,
45
+ :ssim_lambda, :tile_size
46
+ # rubocop:enable Naming/MethodName
47
+
48
+ def initialize(**options)
49
+ unknown = options.keys - DEFAULTS.keys
50
+ raise ArgumentError, "unknown trainer options: #{unknown.join(', ')}" unless unknown.empty?
51
+
52
+ DEFAULTS.merge(options).each do |name, value|
53
+ value = value.dup if value.is_a?(Array)
54
+ instance_variable_set(:"@#{name}", value)
55
+ end
56
+ validate!
57
+ end
58
+
59
+ # Returns an independent hash suitable for checkpoint metadata.
60
+ #
61
+ # @return [Hash{Symbol=>Object}]
62
+ def to_h
63
+ DEFAULTS.keys.to_h { |name| [name, public_send(name)] }
64
+ end
65
+
66
+ private
67
+
68
+ def validate!
69
+ positive_integers = %i[max_steps batch_size sh_degree_interval log_every tile_size]
70
+ invalid = positive_integers.find do |name|
71
+ value = public_send(name)
72
+ !value.is_a?(Integer) || !value.positive?
73
+ end
74
+ raise ArgumentError, "#{invalid} must be a positive integer" if invalid
75
+ raise ArgumentError, "sh_degree must be in 0..4" unless sh_degree.is_a?(Integer) && sh_degree.between?(0, 4)
76
+ raise ArgumentError, "ssim_lambda must be between 0 and 1" unless ssim_lambda.between?(0.0, 1.0)
77
+ unless %i[three_d two_d 3dgs 2dgs].include?(model_type.to_sym)
78
+ raise ArgumentError, "model_type must be :three_d/:3dgs or :two_d/:2dgs"
79
+ end
80
+
81
+ learning_rates = %i[means_lr means_lr_final scales_lr quats_lr opacities_lr sh0_lr shN_lr]
82
+ invalid_rate = learning_rates.find { |name| !public_send(name).positive? }
83
+ raise ArgumentError, "#{invalid_rate} must be positive" if invalid_rate
84
+
85
+ regularizers = %i[opacity_reg scale_reg]
86
+ invalid_regularizer = regularizers.find { |name| public_send(name).negative? }
87
+ raise ArgumentError, "#{invalid_regularizer} must be non-negative" if invalid_regularizer
88
+ end
89
+ end
90
+ end
91
+ end