Class: Gsplat::Strategy::Base

Inherits:
Object
  • Object
show all
Defined in:
lib/gsplat/strategy/base.rb

Overview

Shared strategy lifecycle and parameter/optimizer validation.

Direct Known Subclasses

Default, MCMC

Constant Summary collapse

REQUIRED_KEYS =

Parameter names required by every structural strategy.

%i[means quats scales opacities sh0 shN].freeze

Instance Method Summary collapse

Instance Method Details

#check_sanity(params, optimizers) ⇒ Object

rubocop:disable Metrics/AbcSize, Naming/PredicateMethod

Raises:

  • (ArgumentError)


17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
# File 'lib/gsplat/strategy/base.rb', line 17

def check_sanity(params, optimizers)
  missing_params = REQUIRED_KEYS - params.keys
  missing_optimizers = REQUIRED_KEYS - optimizers.keys
  raise ArgumentError, "missing params: #{missing_params.join(', ')}" unless missing_params.empty?
  raise ArgumentError, "missing optimizers: #{missing_optimizers.join(', ')}" unless missing_optimizers.empty?

  count = nil
  REQUIRED_KEYS.each do |key|
    variable = params.fetch(key)
    raise ArgumentError, "params[:#{key}] must be an Autograd::Variable" unless variable.is_a?(Autograd::Variable)

    count ||= variable.data.shape[0]
    unless variable.data.ndim.positive? && variable.data.shape[0] == count
      raise ShapeError, "all params must share first-axis size #{count}; #{key}=#{variable.data.shape.inspect}"
    end

    optimizer = optimizers.fetch(key)
    unless optimizer.is_a?(Optim::Adam) &&
           optimizer.groups.values.any? { |group| group.variable.equal?(variable) }
      raise ArgumentError, "optimizers[:#{key}] must optimize params[:#{key}]"
    end
  end
  true
end

#initialize_state(scene_scale:) ⇒ Object

Raises:

  • (ArgumentError)


10
11
12
13
14
# File 'lib/gsplat/strategy/base.rb', line 10

def initialize_state(scene_scale:)
  raise ArgumentError, "scene_scale must be positive" unless scene_scale.positive?

  { scene_scale: scene_scale }
end

#step_pre_backward(params:, optimizers:, state:, step:, info:) ⇒ Object

rubocop:enable Metrics/AbcSize, Naming/PredicateMethod



43
44
45
46
# File 'lib/gsplat/strategy/base.rb', line 43

def step_pre_backward(params:, optimizers:, state:, step:, info:)
  [params, optimizers, state, step, info]
  nil
end