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
-
.load(source) ⇒ Snapshot
Loads parameters and optimizer state without mutating live objects.
-
.restore!(source, params:, optimizers:) ⇒ Snapshot
Loads a checkpoint into existing Variables and optimizers.
- .save(target, params:, optimizers:, step:, config: {}) ⇒ Object
Class Method Details
.load(source) ⇒ Snapshot
Loads parameters and optimizer state without mutating live objects.
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.}" end |
.restore!(source, params:, optimizers:) ⇒ Snapshot
Loads a checkpoint into existing Variables and optimizers.
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.}" end |
.save(target, params:, optimizers:, step:, config: {}) ⇒ Object
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 |