cohere-transcribe 0.1.2 → 0.1.3

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.
Files changed (37) hide show
  1. checksums.yaml +4 -4
  2. data/CHANGELOG.md +9 -0
  3. data/README.md +9 -7
  4. data/ext/cohere_transcribe_native/.gitignore +6 -0
  5. data/ext/cohere_transcribe_native/CMakeLists.txt +8 -0
  6. data/ext/cohere_transcribe_native/audio_abi.cpp +20 -0
  7. data/ext/cohere_transcribe_native/audio_exports.macos +1 -0
  8. data/ext/cohere_transcribe_native/audio_exports.map +1 -0
  9. data/ext/cohere_transcribe_native/test/abi_smoke.rb +26 -0
  10. data/lib/cohere/transcribe/alignment/aligner.rb +40 -7
  11. data/lib/cohere/transcribe/asr/native.rb +43 -97
  12. data/lib/cohere/transcribe/audio/decoder.rb +184 -75
  13. data/lib/cohere/transcribe/audio/ffmpeg_native.rb +84 -32
  14. data/lib/cohere/transcribe/audio/segmentation.rb +1 -0
  15. data/lib/cohere/transcribe/cli.rb +30 -22
  16. data/lib/cohere/transcribe/configuration.rb +13 -0
  17. data/lib/cohere/transcribe/constants.rb +1 -1
  18. data/lib/cohere/transcribe/doctor.rb +2 -1
  19. data/lib/cohere/transcribe/hub.rb +446 -58
  20. data/lib/cohere/transcribe/input.rb +18 -3
  21. data/lib/cohere/transcribe/internal/interruptible_native_call.rb +70 -0
  22. data/lib/cohere/transcribe/internal/session_ownership.rb +61 -0
  23. data/lib/cohere/transcribe/internal/utf8.rb +19 -0
  24. data/lib/cohere/transcribe/model_identity.rb +10 -1
  25. data/lib/cohere/transcribe/output/publication.rb +95 -86
  26. data/lib/cohere/transcribe/pytorch_checkpoint.rb +38 -22
  27. data/lib/cohere/transcribe/runtime/engine.rb +37 -16
  28. data/lib/cohere/transcribe/runtime/preparation.rb +212 -35
  29. data/lib/cohere/transcribe/runtime/resources.rb +28 -62
  30. data/lib/cohere/transcribe/state/checkpoint.rb +3 -2
  31. data/lib/cohere/transcribe/state/io.rb +217 -117
  32. data/lib/cohere/transcribe/state/locking.rb +383 -64
  33. data/lib/cohere/transcribe/state/manifest.rb +21 -22
  34. data/lib/cohere/transcribe/types.rb +25 -11
  35. data/lib/cohere/transcribe/version.rb +1 -1
  36. data/sig/cohere/transcribe.rbs +1 -0
  37. metadata +5 -1
@@ -4,6 +4,7 @@ require "find"
4
4
 
5
5
  require_relative "constants"
6
6
  require_relative "errors"
7
+ require_relative "internal/utf8"
7
8
  require_relative "python_text"
8
9
 
9
10
  module Cohere
@@ -31,9 +32,11 @@ module Cohere
31
32
 
32
33
  text = value.is_a?(String) ? value : value.to_path
33
34
  raise TranscriptionInputError, "#{label} must resolve to a text path" unless text.is_a?(String)
35
+
36
+ text = utf8_path(text, label)
34
37
  raise TranscriptionInputError, "#{label} must not be empty" if PythonText.blank?(text)
35
38
 
36
- text.dup.freeze
39
+ text.freeze
37
40
  rescue EncodingError, TypeError, ArgumentError, SystemCallError => e
38
41
  raise e if e.is_a?(TranscriptionInputError)
39
42
 
@@ -83,8 +86,11 @@ module Cohere
83
86
  path.file? && AUDIO_EXTENSIONS.include?(path.extname.downcase)
84
87
  end
85
88
  end
89
+ paths = paths.map { |path| Pathname(utf8_path(path.to_s, "discovered audio")) }
86
90
  paths.sort_by { |path| path.to_s.downcase(:fold) }.map do |path|
87
- [strict_realpath(path), path.relative_path_from(source)]
91
+ relative_path = path.relative_path_from(source)
92
+ relative_text = utf8_path(relative_path.to_s, "discovered audio")
93
+ [strict_realpath(path), Pathname(relative_text)]
88
94
  end
89
95
  rescue EncodingError, SystemCallError, ArgumentError => e
90
96
  raise TranscriptionInputError, "Cannot inspect input #{source}: #{e.message}"
@@ -92,7 +98,8 @@ module Cohere
92
98
  private_class_method :directory_candidates
93
99
 
94
100
  def strict_realpath(value)
95
- Pathname(value).expand_path.realpath
101
+ resolved = Pathname(value).expand_path.realpath
102
+ Pathname(utf8_path(resolved.to_s, "resolved input"))
96
103
  rescue Errno::ENOENT
97
104
  raise TranscriptionInputError, "Input does not exist: #{Pathname(value).expand_path}"
98
105
  rescue EncodingError, SystemCallError, ArgumentError => e
@@ -105,6 +112,14 @@ module Cohere
105
112
  value.is_a?(String) || value.respond_to?(:to_path)
106
113
  end
107
114
  private_class_method :path_like?
115
+
116
+ def utf8_path(value, label)
117
+ text = Internal::UTF8.normalize(value)
118
+ return text if text
119
+
120
+ raise TranscriptionInputError, "#{label} path must contain valid UTF-8"
121
+ end
122
+ private_class_method :utf8_path
108
123
  end
109
124
  end
110
125
  end
@@ -0,0 +1,70 @@
1
+ # frozen_string_literal: true
2
+
3
+ module Cohere
4
+ module Transcribe
5
+ module Internal
6
+ # Runs a blocking foreign call on a dedicated worker so the caller remains
7
+ # interruptible. Callbacks are retained only for this invocation; callers
8
+ # remain responsible for supplying non-poisoning cancellation behavior.
9
+ module InterruptibleNativeCall
10
+ module_function
11
+
12
+ def run(cancel:, join_interval:, missing_outcome:, thread_name: nil)
13
+ raise ArgumentError, "an operation block is required" unless block_given?
14
+
15
+ outcome = Queue.new
16
+ worker = nil
17
+ Thread.handle_interrupt(Object => :on_blocking) do
18
+ worker = Thread.new do
19
+ outcome << [:returned, yield]
20
+ rescue Exception => e # rubocop:disable Lint/RescueException -- transfer native cancellation intact
21
+ outcome << [:raised, e]
22
+ end
23
+ worker.name = thread_name if thread_name && worker.respond_to?(:name=)
24
+ worker.report_on_exception = false
25
+ worker.join
26
+ ensure
27
+ if worker&.alive?
28
+ Thread.handle_interrupt(Object => :never) do
29
+ cancel_and_hard_join(worker, cancel, join_interval)
30
+ end
31
+ end
32
+ end
33
+
34
+ status, value = begin
35
+ outcome.pop(true)
36
+ rescue ThreadError
37
+ raise missing_outcome
38
+ end
39
+ raise value if status == :raised
40
+
41
+ value
42
+ end
43
+
44
+ # Cancellation can arrive before the worker enters the foreign call.
45
+ # Retry until it exits, suppressing every secondary exception so the
46
+ # first caller exception remains intact.
47
+ def cancel_and_hard_join(worker, cancel, join_interval)
48
+ loop do
49
+ begin
50
+ return if worker.join(0)
51
+ rescue Exception # rubocop:disable Lint/RescueException -- preserve the first caller exception
52
+ nil
53
+ end
54
+ begin
55
+ cancel.call
56
+ rescue Exception # rubocop:disable Lint/RescueException -- preserve the first caller exception
57
+ nil
58
+ end
59
+ begin
60
+ return if worker.join(join_interval)
61
+ rescue Exception # rubocop:disable Lint/RescueException -- preserve the first caller exception
62
+ nil
63
+ end
64
+ end
65
+ end
66
+ private_class_method :cancel_and_hard_join
67
+ end
68
+ end
69
+ end
70
+ end
@@ -0,0 +1,61 @@
1
+ # frozen_string_literal: true
2
+
3
+ require_relative "../errors"
4
+
5
+ module Cohere
6
+ module Transcribe
7
+ module Internal
8
+ # Detached, thread-safe session state for explicit close and ObjectSpace
9
+ # finalization. The close callable must not retain the owner whose
10
+ # finalizer holds this object.
11
+ class SessionOwnership
12
+ CLOSE_RUBY_SESSION = :close.to_proc.freeze
13
+
14
+ def self.finalizer_for(ownership)
15
+ proc { |_object_id| ownership.finalize }
16
+ end
17
+
18
+ def initialize(close: CLOSE_RUBY_SESSION, installed_error: "Session ownership is already installed")
19
+ raise ArgumentError, "close must respond to call" unless close.respond_to?(:call)
20
+
21
+ @close = close
22
+ @installed_error = installed_error.freeze
23
+ @mutex = Mutex.new
24
+ @session = nil
25
+ end
26
+
27
+ def session
28
+ @mutex.synchronize { @session }
29
+ end
30
+
31
+ def install(session)
32
+ @mutex.synchronize do
33
+ raise TranscriptionRuntimeError, @installed_error if @session
34
+
35
+ @session = session
36
+ end
37
+ end
38
+
39
+ def close
40
+ Thread.handle_interrupt(Object => :never) do
41
+ session = @mutex.synchronize do
42
+ current = @session
43
+ @session = nil
44
+ current
45
+ end
46
+ return unless session
47
+
48
+ @close.call(session)
49
+ end
50
+ nil
51
+ end
52
+
53
+ def finalize
54
+ close
55
+ rescue Exception # rubocop:disable Lint/RescueException -- finalizers must not escape during GC or shutdown
56
+ nil
57
+ end
58
+ end
59
+ end
60
+ end
61
+ end
@@ -0,0 +1,19 @@
1
+ # frozen_string_literal: true
2
+
3
+ module Cohere
4
+ module Transcribe
5
+ module Internal
6
+ # Normalizes path/reference bytes without applying locale-dependent
7
+ # transcoding. Callers retain control over their public error type and
8
+ # message when the bytes are not valid UTF-8.
9
+ module UTF8
10
+ module_function
11
+
12
+ def normalize(value)
13
+ text = value.b.force_encoding(Encoding::UTF_8)
14
+ text if text.valid_encoding?
15
+ end
16
+ end
17
+ end
18
+ end
19
+ end
@@ -2,6 +2,7 @@
2
2
 
3
3
  require "json"
4
4
  require_relative "constants"
5
+ require_relative "internal/utf8"
5
6
 
6
7
  module Cohere
7
8
  module Transcribe
@@ -42,7 +43,7 @@ module Cohere
42
43
  raise ArgumentError,
43
44
  "Cannot resolve #{description.downcase} path #{reference.inspect}: #{e.message}"
44
45
  end
45
- return path.realpath.to_s if path.directory?
46
+ return utf8_resolved_path(path.realpath.to_s, description) if path.directory?
46
47
  if path.exist? || path.symlink?
47
48
  raise ArgumentError,
48
49
  "#{description} path #{reference.inspect} is not a directory"
@@ -59,6 +60,14 @@ module Cohere
59
60
  raise ArgumentError, "Cannot resolve #{description.downcase} path #{reference.inspect}: #{e.message}"
60
61
  end
61
62
 
63
+ def utf8_resolved_path(value, description)
64
+ text = Internal::UTF8.normalize(value)
65
+ return text if text
66
+
67
+ raise ArgumentError, "#{description} resolved path must contain valid UTF-8"
68
+ end
69
+ private_class_method :utf8_resolved_path
70
+
62
71
  # Path.expanduser() expands only the leading home-directory component;
63
72
  # unlike Pathname#expand_path, it does not lexically erase a missing
64
73
  # component followed by `..` before the filesystem probe.
@@ -269,7 +269,7 @@ module Cohere
269
269
  plan.directory_bindings.each(&:verify!) if plan&.paths&.any?
270
270
  end
271
271
 
272
- def write(plan, result, options, generation_id: nil, speech_spans: nil)
272
+ def write(plan, result, options, generation_id: nil, speech_spans: nil, lock: nil)
273
273
  return [].freeze if plan.paths.empty?
274
274
 
275
275
  contents = plan.paths.to_h do |format, _path|
@@ -291,7 +291,8 @@ module Cohere
291
291
  transaction_paths,
292
292
  transaction_contents,
293
293
  before_commit: -> { State.ensure_source_unchanged!(snapshot) },
294
- directory_bindings: plan.directory_bindings
294
+ directory_bindings: plan.directory_bindings,
295
+ commit_guard: lock&.method(:verify!)
295
296
  )
296
297
  plan.paths.values.freeze
297
298
  end
@@ -844,13 +845,15 @@ module Cohere
844
845
  private_class_method :quantile
845
846
 
846
847
  def atomic_write_set(paths, contents, before_commit: nil, rename: nil,
847
- directory_bindings: nil, operation_hook: nil)
848
+ directory_bindings: nil, operation_hook: nil, commit_guard: nil)
848
849
  staged = {}
849
850
  backups = {}
850
851
  committed = []
851
852
  preserved_backups = {}
852
853
  open_handles = []
853
- failed = false
854
+ completed = false
855
+ primary_error = nil
856
+ recovery_errors = []
854
857
  bound_directories = publication_bound_directories(paths, directory_bindings)
855
858
  operation_hook&.call(:directories_opened, nil)
856
859
  paths.each do |format, supplied_destination|
@@ -900,90 +903,126 @@ module Cohere
900
903
  backups[destination] = [bound, backup_name].freeze
901
904
  open_handles << backup
902
905
  end
906
+ backup_copy_error = nil
903
907
  begin
904
908
  backup.chmod(source.stat.mode & 0o7777)
905
909
  operation_hook&.call(:before_backup_copy, destination)
906
910
  IO.copy_stream(source, backup)
907
911
  backup.flush
908
912
  backup.fsync
913
+ rescue Exception => e # rubocop:disable Lint/RescueException -- retain the copy failure through close
914
+ backup_copy_error = e
915
+ raise
909
916
  ensure
910
- source.close
911
- backup.close
917
+ close_failures = State.close_atomic_resources(
918
+ [source, backup],
919
+ recovery_errors,
920
+ label: ->(handle) { handle.respond_to?(:path) ? handle.path : destination }
921
+ )
922
+ raise close_failures.first if close_failures.any? && backup_copy_error.nil?
912
923
  end
913
924
  bound.verify!
914
925
  end
915
926
  before_commit&.call
916
927
  bound_directories.each_value(&:verify!)
928
+ Thread.handle_interrupt(State::DEFERRED_PUBLICATION_EXCEPTIONS) { commit_guard&.call }
917
929
  staged.each do |destination, (bound, temporary_name)|
918
930
  operation_hook&.call(:before_rename, destination)
919
931
  if rename
932
+ # The caller-supplied operation may block. Record its rollback
933
+ # responsibility first, then leave it interruptible; if it is
934
+ # interrupted before or after moving the entry, the same backup
935
+ # restoration path is safe in both cases.
936
+ Thread.handle_interrupt(State::DEFERRED_PUBLICATION_EXCEPTIONS) { commit_guard&.call }
937
+ committed << destination
938
+ rename.call(bound.display_path(temporary_name), destination)
920
939
  Thread.handle_interrupt(State::DEFERRED_PUBLICATION_EXCEPTIONS) do
921
- committed << destination
922
- rename.call(bound.display_path(temporary_name), destination)
923
940
  bound.rename(temporary_name, destination.basename.to_s) if bound.regular_entry?(temporary_name)
941
+ commit_guard&.call
924
942
  end
925
943
  else
926
944
  Thread.handle_interrupt(State::DEFERRED_PUBLICATION_EXCEPTIONS) do
945
+ commit_guard&.call
927
946
  bound.rename(temporary_name, destination.basename.to_s)
928
947
  committed << destination
929
948
  end
930
949
  end
931
950
  operation_hook&.call(:after_rename, destination)
932
951
  bound.verify!
952
+ Thread.handle_interrupt(State::DEFERRED_PUBLICATION_EXCEPTIONS) { commit_guard&.call }
933
953
  end
934
- bound_directories.each_value(&:fsync)
935
- bound_directories.each_value(&:verify!)
936
- rescue Exception => e # rubocop:disable Lint/RescueException -- rollback must include interrupts
937
- failed = true
938
- rollback_errors = []
939
- committed.reverse_each do |destination|
940
- backup = backups[destination]
941
- bound, backup_name = backup
942
- bound ||= staged.fetch(destination).first
943
- begin
944
- Thread.handle_interrupt(State::DEFERRED_PUBLICATION_EXCEPTIONS) do
945
- bound.unlink(destination.basename.to_s, missing_ok: true)
946
- bound.rename(backup_name, destination.basename.to_s) if backup_name
947
- end
948
- rescue SystemCallError, TranscriptionRuntimeError => rollback_error
949
- preserved_backups[backup] = true if backup_name
950
- rollback_errors << "#{destination}: #{rollback_error.message}"
954
+ bound_directories.each_value do |bound|
955
+ Thread.handle_interrupt(State::DEFERRED_PUBLICATION_EXCEPTIONS) do
956
+ commit_guard&.call
957
+ bound.fsync
958
+ commit_guard&.call
951
959
  end
960
+ bound.verify!
952
961
  end
953
- begin
954
- bound_directories&.each_value(&:fsync)
955
- rescue SystemCallError, TranscriptionRuntimeError => rollback_error
956
- rollback_errors << "directory sync: #{rollback_error.message}"
957
- end
958
- if rollback_errors.any?
959
- detail = rollback_errors.join("; ")
960
- retained = preserved_backups.keys.filter_map do |backup|
961
- backup&.then { |bound, name| bound.display_path(name).to_s }
962
- end.sort
963
- raise TranscriptionRuntimeError,
964
- "Output commit failed and rollback was incomplete (#{detail}); " \
965
- "preserved backups: #{retained}",
966
- cause: e
962
+ Thread.handle_interrupt(State::DEFERRED_PUBLICATION_EXCEPTIONS) do
963
+ commit_guard&.call
964
+ completed = true
967
965
  end
968
- raise e
966
+ rescue Exception => e # rubocop:disable Lint/RescueException -- rollback must include interrupts
967
+ primary_error = e
968
+ raise
969
969
  ensure
970
- cleanup_errors = []
971
- staged&.each_value do |bound, name|
972
- bound.unlink(name, missing_ok: true)
973
- rescue SystemCallError, TranscriptionRuntimeError => e
974
- cleanup_errors << e
975
- end
976
- backups&.each_value do |backup|
977
- next unless backup && !preserved_backups&.key?(backup)
978
-
979
- bound, name = backup
980
- bound.unlink(name, missing_ok: true)
981
- rescue SystemCallError, TranscriptionRuntimeError => e
982
- cleanup_errors << e
970
+ Thread.handle_interrupt(State::DEFERRED_PUBLICATION_EXCEPTIONS) do
971
+ unless completed
972
+ committed.reverse_each do |destination|
973
+ backup = backups[destination]
974
+ bound, backup_name = backup
975
+ bound ||= staged.fetch(destination).first
976
+ begin
977
+ bound.unlink(destination.basename.to_s, missing_ok: true)
978
+ bound.rename(backup_name, destination.basename.to_s) if backup_name
979
+ backups[destination] = nil
980
+ rescue SystemCallError, TranscriptionRuntimeError => e
981
+ preserved_backups[backup] = true if backup_name
982
+ recovery_errors << "#{destination}: #{e.message}"
983
+ end
984
+ end
985
+ begin
986
+ bound_directories&.each_value(&:fsync)
987
+ rescue SystemCallError, TranscriptionRuntimeError => e
988
+ recovery_errors << "directory sync: #{e.message}"
989
+ end
990
+ end
991
+ staged&.each_value do |bound, name|
992
+ bound.unlink(name, missing_ok: true)
993
+ rescue SystemCallError, TranscriptionRuntimeError => e
994
+ recovery_errors << "cleanup #{bound.display_path(name)}: #{e.message}"
995
+ end
996
+ backups&.each_value do |backup|
997
+ next unless backup && !preserved_backups&.key?(backup)
998
+
999
+ bound, name = backup
1000
+ bound.unlink(name, missing_ok: true)
1001
+ rescue SystemCallError, TranscriptionRuntimeError => e
1002
+ preserved_backups[backup] = true
1003
+ recovery_errors << "cleanup #{bound.display_path(name)}: #{e.message}"
1004
+ end
1005
+ State.close_atomic_resources(
1006
+ open_handles,
1007
+ recovery_errors,
1008
+ label: ->(handle) { handle.respond_to?(:path) ? handle.path : handle.inspect }
1009
+ )
1010
+ State.close_atomic_resources(
1011
+ bound_directories&.values,
1012
+ recovery_errors,
1013
+ label: ->(bound) { bound.binding.canonical_path }
1014
+ )
1015
+ retained = preserved_backups.keys.filter_map do |backup|
1016
+ backup&.then { |bound, name| bound.display_path(name) }
1017
+ end
1018
+ State.finish_atomic_recovery!(
1019
+ subject: "Output commit",
1020
+ completed: completed,
1021
+ primary_error: primary_error,
1022
+ recovery_errors: recovery_errors,
1023
+ retained_backups: retained
1024
+ )
983
1025
  end
984
- open_handles&.each { |handle| handle.close unless handle.closed? }
985
- bound_directories&.each_value(&:close)
986
- raise cleanup_errors.first if cleanup_errors.any? && !failed
987
1026
  end
988
1027
 
989
1028
  def publication_bound_directories(paths, directory_bindings)
@@ -1021,36 +1060,6 @@ module Cohere
1021
1060
  end
1022
1061
  private_class_method :bound_output_mode
1023
1062
 
1024
- def output_mode(destination)
1025
- return destination.stat.mode & 0o7777 if destination.exist?
1026
-
1027
- current_umask = File.umask
1028
- File.umask(current_umask)
1029
- 0o666 & ~current_umask
1030
- end
1031
- private_class_method :output_mode
1032
-
1033
- def fsync_directories(directories)
1034
- directories.each do |directory|
1035
- File.open(directory, File::RDONLY, &:fsync)
1036
- rescue Errno::EACCES, Errno::EBADF, Errno::EINVAL, Errno::EISDIR,
1037
- Errno::ENOTSUP, Errno::EPERM
1038
- next
1039
- end
1040
- end
1041
- private_class_method :fsync_directories
1042
-
1043
- def verified_publication?(source_snapshot, paths, state_path, options)
1044
- State.verify_published_outputs(
1045
- source_snapshot: source_snapshot,
1046
- output_paths: paths,
1047
- state_path: state_path,
1048
- asr_contract_key: State.asr_contract_key(options),
1049
- render_contract_key: State.render_contract_key(options)
1050
- ).verified?
1051
- end
1052
- private_class_method :verified_publication?
1053
-
1054
1063
  def source_record(path)
1055
1064
  State::SourceSnapshot.capture(path)
1056
1065
  end
@@ -184,8 +184,8 @@ module Cohere
184
184
  total_entries == ZIP64_UINT16_MARKER ||
185
185
  central_size == ZIP64_UINT32_MARKER ||
186
186
  central_offset == ZIP64_UINT32_MARKER
187
- total_entries, central_size, central_offset = zip64_directory(eocd_offset) if
188
- zip64_marker && zip64_locator_present?(eocd_offset)
187
+ zip64_values = zip64_directory(eocd_offset) if zip64_marker
188
+ total_entries, central_size, central_offset = zip64_values if zip64_values
189
189
  if total_entries > ZIP_ENTRY_LIMIT || central_size > ZIP_CENTRAL_LIMIT ||
190
190
  central_size > size || central_offset > size - central_size
191
191
  raise Error, "PyTorch archive #{path} has invalid central-directory bounds"
@@ -204,11 +204,12 @@ module Cohere
204
204
  end
205
205
 
206
206
  def zip64_directory(eocd_offset)
207
- raise Error, "PyTorch archive #{path} has no ZIP64 locator" if eocd_offset < 20
207
+ return if eocd_offset < 20
208
208
 
209
209
  locator = read_range(eocd_offset - 20, 20)
210
210
  signature, disk, record_offset, disks = locator.unpack("VVQ<V")
211
- raise Error, "PyTorch archive #{path} has an invalid ZIP64 locator" unless signature == ZIP64_LOCATOR && disk.zero? && disks == 1
211
+ return unless signature == ZIP64_LOCATOR
212
+ raise Error, "PyTorch archive #{path} has an invalid ZIP64 locator" unless disk.zero? && disks == 1
212
213
 
213
214
  record = read_range(record_offset, 56)
214
215
  signature, record_size, _made, _needed, disk, central_disk,
@@ -221,12 +222,6 @@ module Cohere
221
222
  [total_entries, central_size, central_offset]
222
223
  end
223
224
 
224
- def zip64_locator_present?(eocd_offset)
225
- return false if eocd_offset < 20
226
-
227
- read_range(eocd_offset - 20, 4).unpack1("V") == ZIP64_LOCATOR
228
- end
229
-
230
225
  def parse_central_directory(bytes, expected_entries:, archive_size:, data_limit:)
231
226
  offset = 0
232
227
  result = {}
@@ -783,7 +778,6 @@ module Cohere
783
778
  @path = Pathname(path).expand_path
784
779
  @storages = {}
785
780
  @storage_payloads = {}
786
- @verified_storage_payloads = {}
787
781
  @storage_verification_mutex = Mutex.new
788
782
  @tensors = read_checkpoint.freeze
789
783
  end
@@ -812,20 +806,22 @@ module Cohere
812
806
  end
813
807
  raise ArgumentError, "chunk_bytes must be positive" unless chunk_bytes.positive?
814
808
 
815
- verify_storage_payload!(tensor)
816
-
817
809
  source_width = Safetensors::DTYPE_BYTES.fetch(tensor.dtype)
818
810
  elements_per_chunk = [chunk_bytes / source_width, 1].max
819
811
  bytes_per_chunk = elements_per_chunk * source_width
820
812
  unless contiguous?(tensor)
813
+ verify_storage_payload!(tensor)
821
814
  return write_strided_tensor(
822
815
  tensor, output, target_dtype: target_dtype, converter: converter,
823
816
  elements_per_chunk: elements_per_chunk, source_width: source_width
824
817
  )
825
818
  end
826
819
 
820
+ inline_payload = inline_storage_verification(tensor)
821
+ verify_storage_payload!(tensor) unless inline_payload
827
822
  remaining = tensor.nbytes
828
823
  written = 0
824
+ crc32 = 0 if inline_payload
829
825
  File.open(path, "rb") do |source|
830
826
  source.seek(tensor.data_start)
831
827
  while remaining.positive?
@@ -833,16 +829,18 @@ module Cohere
833
829
  chunk = source.read(requested)
834
830
  raise Error, "Unexpected end of #{path} while reading #{tensor.name.inspect}" if chunk.nil? || chunk.bytesize != requested
835
831
 
832
+ crc32 = Zlib.crc32(chunk, crc32) if inline_payload
836
833
  converted = converter.convert(chunk, from: tensor.dtype, to: target_dtype)
837
834
  output.write(converted)
838
835
  written += converted.bytesize
839
836
  remaining -= requested
840
837
  end
841
838
  end
839
+ mark_inline_storage_verified!(tensor.storage_key, inline_payload, crc32) if inline_payload
842
840
  expected = tensor.element_count * Safetensors::DTYPE_BYTES.fetch(target_dtype)
843
- return written if written == expected
841
+ raise Error, "Converted tensor #{tensor.name.inspect} wrote #{written} bytes; expected #{expected}" unless written == expected
844
842
 
845
- raise Error, "Converted tensor #{tensor.name.inspect} wrote #{written} bytes; expected #{expected}"
843
+ written
846
844
  end
847
845
 
848
846
  private
@@ -1036,14 +1034,28 @@ module Cohere
1036
1034
  end
1037
1035
 
1038
1036
  def verify_storage_payload!(tensor)
1039
- payload = @storage_payloads[tensor.storage_key]
1040
- return unless payload
1041
-
1042
1037
  @storage_verification_mutex.synchronize do
1043
- return if @verified_storage_payloads.key?(tensor.storage_key)
1038
+ payload = @storage_payloads[tensor.storage_key]
1039
+ return unless payload
1044
1040
 
1045
1041
  verify_storage_crc!(tensor.storage_key, payload)
1046
- @verified_storage_payloads[tensor.storage_key] = true
1042
+ @storage_payloads.delete(tensor.storage_key)
1043
+ end
1044
+ end
1045
+
1046
+ def inline_storage_verification(tensor)
1047
+ payload = @storage_verification_mutex.synchronize do
1048
+ @storage_payloads[tensor.storage_key]
1049
+ end
1050
+ return unless payload && tensor.data_start == payload.data_offset && tensor.nbytes == payload.size
1051
+
1052
+ payload
1053
+ end
1054
+
1055
+ def mark_inline_storage_verified!(storage_key, payload, crc32)
1056
+ verify_storage_crc_value!(storage_key, payload, crc32)
1057
+ @storage_verification_mutex.synchronize do
1058
+ @storage_payloads.delete(storage_key) if @storage_payloads[storage_key].equal?(payload)
1047
1059
  end
1048
1060
  end
1049
1061
 
@@ -1063,11 +1075,15 @@ module Cohere
1063
1075
  remaining -= requested
1064
1076
  end
1065
1077
  end
1078
+ verify_storage_crc_value!(storage_key, payload, crc32)
1079
+ rescue Errno::ENOENT, Errno::EACCES => e
1080
+ raise Error, "Cannot read PyTorch checkpoint #{path}: #{e.message}"
1081
+ end
1082
+
1083
+ def verify_storage_crc_value!(storage_key, payload, crc32)
1066
1084
  return if crc32 == payload.crc32
1067
1085
 
1068
1086
  raise Error, "PyTorch tensor storage #{storage_key.inspect} failed its CRC check"
1069
- rescue Errno::ENOENT, Errno::EACCES => e
1070
- raise Error, "Cannot read PyTorch checkpoint #{path}: #{e.message}"
1071
1087
  end
1072
1088
 
1073
1089
  # General strided tensors occur in legitimate torch state dictionaries