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.
- checksums.yaml +7 -0
- data/LICENSE.txt +202 -0
- data/README.md +236 -0
- data/docs/ACCEPTANCE.md +60 -0
- data/docs/BENCHMARKS.md +79 -0
- data/docs/DECISIONS.md +52 -0
- data/docs/MIGRATION.md +117 -0
- data/docs/PROFILE.md +75 -0
- data/docs/PROGRESS.md +199 -0
- data/docs/decisions/0000-template.md +17 -0
- data/docs/decisions/0001-pin-golden-reference-and-cpu-validation.md +27 -0
- data/docs/decisions/0002-retain-ruby-fallbacks-for-native-backend.md +26 -0
- data/docs/decisions/0003-share-projection-for-distorted-cameras.md +26 -0
- data/docs/decisions/0004-use-portable-world-space-reference-paths.md +26 -0
- data/docs/decisions/0005-reuse-ewa-core-for-2dgs.md +24 -0
- data/docs/decisions/0006-keep-eval3d-as-portable-reference.md +25 -0
- data/docs/decisions/0007-share-compositor-semantics-for-contribution-indices.md +25 -0
- data/examples/data/README.md +14 -0
- data/examples/data/colmap/images/view_000.png +0 -0
- data/examples/data/colmap/images/view_001.png +0 -0
- data/examples/data/colmap/images/view_002.png +0 -0
- data/examples/data/colmap/sparse/0/cameras.txt +2 -0
- data/examples/data/colmap/sparse/0/images.txt +6 -0
- data/examples/data/colmap/sparse/0/points3D.txt +17 -0
- data/examples/data/splats.ply +0 -0
- data/examples/fit_image.rb +35 -0
- data/examples/generate_sample_data.rb +147 -0
- data/examples/render_path.rb +117 -0
- data/examples/simple_trainer.rb +76 -0
- data/ext/gsplat_native/common.h +54 -0
- data/ext/gsplat_native/extconf.rb +39 -0
- data/ext/gsplat_native/gsplat_native.c +60 -0
- data/ext/gsplat_native/intersections.c +211 -0
- data/ext/gsplat_native/projection.c +199 -0
- data/ext/gsplat_native/raster_backward.c +129 -0
- data/ext/gsplat_native/raster_backward_bridge.c +69 -0
- data/ext/gsplat_native/raster_forward.c +150 -0
- data/ext/gsplat_native/rasterization.h +51 -0
- data/ext/gsplat_native/spherical_harmonics.c +129 -0
- data/gsplat.gemspec +36 -0
- data/lib/gsplat/autograd/context.rb +60 -0
- data/lib/gsplat/autograd/function.rb +68 -0
- data/lib/gsplat/autograd/variable.rb +159 -0
- data/lib/gsplat/backend/ruby/accumulate.rb +139 -0
- data/lib/gsplat/backend/ruby/accumulate_backward.rb +40 -0
- data/lib/gsplat/backend/ruby/eval3d_rasterizer.rb +175 -0
- data/lib/gsplat/backend/ruby/isect_tiles.rb +198 -0
- data/lib/gsplat/backend/ruby/projection.rb +251 -0
- data/lib/gsplat/backend/ruby/projection_backward.rb +190 -0
- data/lib/gsplat/backend/ruby/projection_covariance_vjp.rb +72 -0
- data/lib/gsplat/backend/ruby/projection_input_vjp.rb +252 -0
- data/lib/gsplat/backend/ruby/quat_scale_to_covar_preci.rb +139 -0
- data/lib/gsplat/backend/ruby/rasterize_to_indices_in_range.rb +121 -0
- data/lib/gsplat/backend/ruby/rasterize_to_pixels.rb +199 -0
- data/lib/gsplat/backend/ruby/rasterize_to_pixels_backward.rb +121 -0
- data/lib/gsplat/backend/ruby/spherical_harmonics.rb +135 -0
- data/lib/gsplat/backend/ruby/tile_compositor.rb +50 -0
- data/lib/gsplat/backend/ruby/tile_compositor_backward.rb +84 -0
- data/lib/gsplat/backend.rb +79 -0
- data/lib/gsplat/compression/grid_sort.rb +79 -0
- data/lib/gsplat/compression/kmeans.rb +121 -0
- data/lib/gsplat/compression/png.rb +147 -0
- data/lib/gsplat/compression/png_codec.rb +133 -0
- data/lib/gsplat/compression/quantizer.rb +79 -0
- data/lib/gsplat/io/checkpoint.rb +145 -0
- data/lib/gsplat/io/colmap.rb +175 -0
- data/lib/gsplat/io/colmap_binary.rb +98 -0
- data/lib/gsplat/io/colmap_text.rb +84 -0
- data/lib/gsplat/io/image.rb +63 -0
- data/lib/gsplat/io/image_backends.rb +81 -0
- data/lib/gsplat/io/npy.rb +189 -0
- data/lib/gsplat/io/ply.rb +185 -0
- data/lib/gsplat/io/ply_reader.rb +142 -0
- data/lib/gsplat/io/zip_archive.rb +183 -0
- data/lib/gsplat/math/camera_distortion.rb +123 -0
- data/lib/gsplat/math/camera_projection.rb +202 -0
- data/lib/gsplat/math/mat.rb +114 -0
- data/lib/gsplat/math/quaternion.rb +175 -0
- data/lib/gsplat/math/small_matrix_primitives.rb +112 -0
- data/lib/gsplat/math/spherical_harmonic_basis.rb +148 -0
- data/lib/gsplat/math/ssim.rb +153 -0
- data/lib/gsplat/native.rb +30 -0
- data/lib/gsplat/native_ops.rb +148 -0
- data/lib/gsplat/native_raster_ops.rb +76 -0
- data/lib/gsplat/ops/accumulate.rb +68 -0
- data/lib/gsplat/ops/eval3d_rasterize.rb +62 -0
- data/lib/gsplat/ops/isect_tiles.rb +38 -0
- data/lib/gsplat/ops/projection.rb +154 -0
- data/lib/gsplat/ops/quat_scale_to_covar_preci.rb +118 -0
- data/lib/gsplat/ops/rasterize_to_indices_in_range.rb +38 -0
- data/lib/gsplat/ops/rasterize_to_pixels.rb +105 -0
- data/lib/gsplat/ops/relocation.rb +94 -0
- data/lib/gsplat/ops/spherical_harmonics.rb +71 -0
- data/lib/gsplat/ops/tensor_shape_ops.rb +173 -0
- data/lib/gsplat/ops/tensor_value_ops.rb +150 -0
- data/lib/gsplat/optim/adam.rb +212 -0
- data/lib/gsplat/optim/lr_scheduler.rb +36 -0
- data/lib/gsplat/optim/selective_adam.rb +68 -0
- data/lib/gsplat/rasterization.rb +130 -0
- data/lib/gsplat/rasterization_2dgs.rb +140 -0
- data/lib/gsplat/rasterization_helpers.rb +162 -0
- data/lib/gsplat/rasterization_validation.rb +99 -0
- data/lib/gsplat/strategy/base.rb +49 -0
- data/lib/gsplat/strategy/default.rb +188 -0
- data/lib/gsplat/strategy/mcmc.rb +103 -0
- data/lib/gsplat/strategy/mcmc_ops.rb +143 -0
- data/lib/gsplat/strategy/ops.rb +165 -0
- data/lib/gsplat/training/config.rb +91 -0
- data/lib/gsplat/training/image_fitter.rb +158 -0
- data/lib/gsplat/training/losses.rb +191 -0
- data/lib/gsplat/training/scene.rb +116 -0
- data/lib/gsplat/training/trainer.rb +236 -0
- data/lib/gsplat/utils.rb +110 -0
- data/lib/gsplat/version.rb +6 -0
- data/lib/gsplat.rb +98 -0
- 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
|