Class: Gsplat::Strategy::Base
- Inherits:
-
Object
- Object
- Gsplat::Strategy::Base
- Defined in:
- lib/gsplat/strategy/base.rb
Overview
Shared strategy lifecycle and parameter/optimizer validation.
Constant Summary collapse
- REQUIRED_KEYS =
Parameter names required by every structural strategy.
%i[means quats scales opacities sh0 shN].freeze
Instance Method Summary collapse
-
#check_sanity(params, optimizers) ⇒ Object
rubocop:disable Metrics/AbcSize, Naming/PredicateMethod.
- #initialize_state(scene_scale:) ⇒ Object
-
#step_pre_backward(params:, optimizers:, state:, step:, info:) ⇒ Object
rubocop:enable Metrics/AbcSize, Naming/PredicateMethod.
Instance Method Details
#check_sanity(params, optimizers) ⇒ Object
rubocop:disable Metrics/AbcSize, Naming/PredicateMethod
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
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 |