Module: Gsplat::IO::Checkpoint

Defined in:
lib/gsplat/io/checkpoint.rb

Overview

Portable NPZ checkpoint for parameters, Adam moments, step, and config.

Defined Under Namespace

Classes: Snapshot

Constant Summary collapse

VERSION =

Current checkpoint schema version.

1
SEPARATOR =

Separator reserved for structured NPZ entry names.

"___"

Class Method Summary collapse

Class Method Details

.load(source) ⇒ Snapshot

Loads parameters and optimizer state without mutating live objects.

Parameters:

  • source (String, #read)

    NPZ checkpoint path or IO

Returns:



54
55
56
57
58
59
60
61
62
63
64
65
66
67
# File 'lib/gsplat/io/checkpoint.rb', line 54

def load(source)
  arrays = Npy.read_npz(source)
  version = arrays.fetch("checkpoint_version")[0].to_i
  raise NotSupportedError, "unsupported checkpoint version #{version}" unless version == VERSION

  Snapshot.new(
    params: extract_params(arrays),
    optimizer_states: extract_optimizer_states(arrays),
    step: arrays.fetch("checkpoint_step")[0].to_i,
    config: decode_json(arrays.fetch("checkpoint_config_json"))
  )
rescue KeyError => e
  raise Gsplat::Error, "invalid checkpoint: #{e.message}"
end

.restore!(source, params:, optimizers:) ⇒ Snapshot

Loads a checkpoint into existing Variables and optimizers.

Parameters:

Returns:



75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
# File 'lib/gsplat/io/checkpoint.rb', line 75

def restore!(source, params:, optimizers:)
  snapshot = load(source)
  snapshot.params.each do |name, value|
    params.fetch(name).replace_data!(value.dup)
  end
  snapshot.optimizer_states.each do |optimizer_name, groups|
    optimizer = optimizers.fetch(optimizer_name)
    groups.each do |group_name, state|
      optimizer.load_state!(
        group_name,
        step: state.fetch(:step),
        exp_avg: state.fetch(:exp_avg),
        exp_avg_sq: state.fetch(:exp_avg_sq)
      )
    end
  end
  snapshot
rescue KeyError => e
  raise Gsplat::Error, "checkpoint target mismatch: #{e.message}"
end

.save(target, params:, optimizers:, step:, config: {}) ⇒ Object

Raises:

  • (ArgumentError)


27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
# File 'lib/gsplat/io/checkpoint.rb', line 27

def save(target, params:, optimizers:, step:, config: {})
  raise ArgumentError, "step must be a non-negative integer" unless step.is_a?(Integer) && !step.negative?

  arrays = {
    "checkpoint_version" => Numo::Int32[VERSION],
    "checkpoint_step" => Numo::Int64[step],
    "checkpoint_config_json" => encode_json(config)
  }
  params.each do |name, value|
    arrays[key("param", name)] = Ops::TensorOps.data(value)
  end
  optimizers.each do |optimizer_name, optimizer|
    optimizer.groups.each_key do |group_name|
      state = optimizer.state(group_name)
      prefix = key("optimizer", optimizer_name, group_name)
      arrays["#{prefix}#{SEPARATOR}step"] = Numo::Int64[state.step]
      arrays["#{prefix}#{SEPARATOR}exp_avg"] = state.exp_avg
      arrays["#{prefix}#{SEPARATOR}exp_avg_sq"] = state.exp_avg_sq
    end
  end
  Npy.write_npz(target, arrays)
end