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.
- checksums.yaml +4 -4
- data/CHANGELOG.md +9 -0
- data/README.md +9 -7
- data/ext/cohere_transcribe_native/.gitignore +6 -0
- data/ext/cohere_transcribe_native/CMakeLists.txt +8 -0
- data/ext/cohere_transcribe_native/audio_abi.cpp +20 -0
- data/ext/cohere_transcribe_native/audio_exports.macos +1 -0
- data/ext/cohere_transcribe_native/audio_exports.map +1 -0
- data/ext/cohere_transcribe_native/test/abi_smoke.rb +26 -0
- data/lib/cohere/transcribe/alignment/aligner.rb +40 -7
- data/lib/cohere/transcribe/asr/native.rb +43 -97
- data/lib/cohere/transcribe/audio/decoder.rb +184 -75
- data/lib/cohere/transcribe/audio/ffmpeg_native.rb +84 -32
- data/lib/cohere/transcribe/audio/segmentation.rb +1 -0
- data/lib/cohere/transcribe/cli.rb +30 -22
- data/lib/cohere/transcribe/configuration.rb +13 -0
- data/lib/cohere/transcribe/constants.rb +1 -1
- data/lib/cohere/transcribe/doctor.rb +2 -1
- data/lib/cohere/transcribe/hub.rb +446 -58
- data/lib/cohere/transcribe/input.rb +18 -3
- data/lib/cohere/transcribe/internal/interruptible_native_call.rb +70 -0
- data/lib/cohere/transcribe/internal/session_ownership.rb +61 -0
- data/lib/cohere/transcribe/internal/utf8.rb +19 -0
- data/lib/cohere/transcribe/model_identity.rb +10 -1
- data/lib/cohere/transcribe/output/publication.rb +95 -86
- data/lib/cohere/transcribe/pytorch_checkpoint.rb +38 -22
- data/lib/cohere/transcribe/runtime/engine.rb +37 -16
- data/lib/cohere/transcribe/runtime/preparation.rb +212 -35
- data/lib/cohere/transcribe/runtime/resources.rb +28 -62
- data/lib/cohere/transcribe/state/checkpoint.rb +3 -2
- data/lib/cohere/transcribe/state/io.rb +217 -117
- data/lib/cohere/transcribe/state/locking.rb +383 -64
- data/lib/cohere/transcribe/state/manifest.rb +21 -22
- data/lib/cohere/transcribe/types.rb +25 -11
- data/lib/cohere/transcribe/version.rb +1 -1
- data/sig/cohere/transcribe.rbs +1 -0
- 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.
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
911
|
-
|
|
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
|
|
935
|
-
|
|
936
|
-
|
|
937
|
-
|
|
938
|
-
|
|
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
|
-
|
|
954
|
-
|
|
955
|
-
|
|
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
|
-
|
|
966
|
+
rescue Exception => e # rubocop:disable Lint/RescueException -- rollback must include interrupts
|
|
967
|
+
primary_error = e
|
|
968
|
+
raise
|
|
969
969
|
ensure
|
|
970
|
-
|
|
971
|
-
|
|
972
|
-
|
|
973
|
-
|
|
974
|
-
|
|
975
|
-
|
|
976
|
-
|
|
977
|
-
|
|
978
|
-
|
|
979
|
-
|
|
980
|
-
|
|
981
|
-
|
|
982
|
-
|
|
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
|
-
|
|
188
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
841
|
+
raise Error, "Converted tensor #{tensor.name.inspect} wrote #{written} bytes; expected #{expected}" unless written == expected
|
|
844
842
|
|
|
845
|
-
|
|
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
|
-
|
|
1038
|
+
payload = @storage_payloads[tensor.storage_key]
|
|
1039
|
+
return unless payload
|
|
1044
1040
|
|
|
1045
1041
|
verify_storage_crc!(tensor.storage_key, payload)
|
|
1046
|
-
@
|
|
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
|