lfeats 0.1.4__tar.gz → 0.2.0__tar.gz

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 (185) hide show
  1. {lfeats-0.1.4 → lfeats-0.2.0}/PKG-INFO +2 -1
  2. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/interfaces/extractor.py +12 -0
  3. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/interfaces/resampler.py +11 -0
  4. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/base.py +19 -0
  5. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/contentvec.py +0 -2
  6. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/data2vec.py +0 -3
  7. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/data2vec2.py +0 -3
  8. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/ecapa_tdnn.py +1 -3
  9. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/emotion2vec.py +1 -3
  10. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/emotion2vec_plus.py +2 -4
  11. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/hubert.py +0 -2
  12. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/manager.py +13 -0
  13. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/next_tdnn.py +0 -2
  14. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/r_spin.py +0 -2
  15. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/r_vector.py +1 -3
  16. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/spidr.py +4 -8
  17. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/spin.py +3 -4
  18. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/sslzip.py +14 -1
  19. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/unispeech_sat.py +0 -2
  20. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/wav2vec2.py +0 -3
  21. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/wavlm.py +0 -2
  22. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/whisper.py +0 -1
  23. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/x_vector.py +1 -3
  24. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/resamplers/base.py +11 -0
  25. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/resamplers/lilfilter.py +13 -1
  26. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/resamplers/manager.py +13 -0
  27. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/resamplers/torchaudio.py +14 -1
  28. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/util/download.py +24 -10
  29. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/utils/distributed.py +1 -1
  30. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/utils/io.py +45 -6
  31. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/version.py +1 -1
  32. {lfeats-0.1.4 → lfeats-0.2.0}/pyproject.toml +1 -0
  33. {lfeats-0.1.4 → lfeats-0.2.0}/.gitignore +0 -0
  34. {lfeats-0.1.4 → lfeats-0.2.0}/LICENSE +0 -0
  35. {lfeats-0.1.4 → lfeats-0.2.0}/README.md +0 -0
  36. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/__init__.py +0 -0
  37. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/cli.py +0 -0
  38. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/interfaces/__init__.py +0 -0
  39. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/interfaces/types.py +0 -0
  40. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/interfaces/utils.py +0 -0
  41. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/models/__init__.py +0 -0
  42. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/resamplers/__init__.py +0 -0
  43. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/resamplers/soxr.py +0 -0
  44. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/__init__.py +0 -0
  45. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/LICENSE +0 -0
  46. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/__init__.py +0 -0
  47. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/checkpoint_utils.py +0 -0
  48. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/config/__init__.py +0 -0
  49. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/config/config.yaml +0 -0
  50. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/data/__init__.py +0 -0
  51. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/data/dictionary.py +0 -0
  52. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/data/modality.py +0 -0
  53. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/data/text_compressor.py +0 -0
  54. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/dataclass/__init__.py +0 -0
  55. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/dataclass/configs.py +0 -0
  56. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/dataclass/constants.py +0 -0
  57. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/dataclass/initialize.py +0 -0
  58. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/dataclass/utils.py +0 -0
  59. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/file_io.py +0 -0
  60. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/incremental_decoding_utils.py +0 -0
  61. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/logging/__init__.py +0 -0
  62. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/logging/meters.py +0 -0
  63. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/__init__.py +0 -0
  64. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/data2vec/__init__.py +0 -0
  65. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/data2vec/data2vec2.py +0 -0
  66. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/data2vec/data2vec_audio.py +0 -0
  67. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/data2vec/modalities/__init__.py +0 -0
  68. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/data2vec/modalities/audio.py +0 -0
  69. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/data2vec/modalities/base.py +0 -0
  70. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/data2vec/modalities/modules.py +0 -0
  71. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/fairseq_decoder.py +0 -0
  72. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/fairseq_encoder.py +0 -0
  73. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/fairseq_incremental_decoder.py +0 -0
  74. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/fairseq_model.py +0 -0
  75. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/hubert/__init__.py +0 -0
  76. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/hubert/hubert.py +0 -0
  77. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/wav2vec/__init__.py +0 -0
  78. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/wav2vec/utils.py +0 -0
  79. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/models/wav2vec/wav2vec2.py +0 -0
  80. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/modules/__init__.py +0 -0
  81. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/modules/ema_module.py +0 -0
  82. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/modules/fairseq_dropout.py +0 -0
  83. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/modules/fp32_group_norm.py +0 -0
  84. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/modules/gelu.py +0 -0
  85. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/modules/gumbel_vector_quantizer.py +0 -0
  86. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/modules/layer_norm.py +0 -0
  87. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/modules/multihead_attention.py +0 -0
  88. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/modules/quant_noise.py +0 -0
  89. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/modules/same_pad.py +0 -0
  90. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/modules/transpose_last.py +0 -0
  91. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/quantization_utils.py +0 -0
  92. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/registry.py +0 -0
  93. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/tasks/__init__.py +0 -0
  94. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/tasks/audio_pretraining.py +0 -0
  95. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/tasks/fairseq_task.py +0 -0
  96. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/tasks/hubert_pretraining.py +0 -0
  97. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/tokenizer.py +0 -0
  98. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/fairseq/utils.py +0 -0
  99. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/LICENSE +0 -0
  100. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/SpeakerNet.py +0 -0
  101. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/__init__.py +0 -0
  102. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/aggregation/__init__.py +0 -0
  103. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/aggregation/vap_bn_tanh_fc_bn.py +0 -0
  104. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/configs/NeXt_TDNN_C256_B3_K65_7.py +0 -0
  105. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/configs/NeXt_TDNN_C256_B3_K65_7_cyclical_lr_step.py +0 -0
  106. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/configs/NeXt_TDNN_light_C256_B3_K65.py +0 -0
  107. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/configs/__init__.py +0 -0
  108. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/main.py +0 -0
  109. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/models/NeXt_TDNN.py +0 -0
  110. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/models/TSConvNeXt.py +0 -0
  111. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/models/TSConvNeXt_light.py +0 -0
  112. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/models/__init__.py +0 -0
  113. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/models/utils.py +0 -0
  114. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/preprocessing/__init__.py +0 -0
  115. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/next_tdnn_asv/preprocessing/mel_transform.py +0 -0
  116. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/rspin/LICENSE +0 -0
  117. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/rspin/__init__.py +0 -0
  118. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/rspin/model.py +0 -0
  119. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/rspin/wavlm_config.py +0 -0
  120. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/LICENSE +0 -0
  121. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/__init__.py +0 -0
  122. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/upstream/__init__.py +0 -0
  123. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/upstream/hubert/__init__.py +0 -0
  124. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/upstream/hubert/convert.py +0 -0
  125. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/upstream/hubert/hubert_model.py +0 -0
  126. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/upstream/utils.py +0 -0
  127. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/upstream/wav2vec2/__init__.py +0 -0
  128. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/upstream/wav2vec2/wav2vec2_model.py +0 -0
  129. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/upstream/wavlm/WavLM.py +0 -0
  130. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/upstream/wavlm/__init__.py +0 -0
  131. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/upstream/wavlm/modules.py +0 -0
  132. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/s3prl/util/__init__.py +0 -0
  133. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/LICENSE +0 -0
  134. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/__init__.py +0 -0
  135. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/dataio/__init__.py +0 -0
  136. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/dataio/dataio.py +0 -0
  137. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/dataio/encoder.py +0 -0
  138. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/dataio/preprocess.py +0 -0
  139. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/inference/__init__.py +0 -0
  140. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/inference/classifiers.py +0 -0
  141. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/inference/interfaces.py +0 -0
  142. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/lobes/__init__.py +0 -0
  143. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/lobes/features.py +0 -0
  144. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/lobes/models/ECAPA_TDNN.py +0 -0
  145. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/lobes/models/ResNet.py +0 -0
  146. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/lobes/models/Xvector.py +0 -0
  147. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/lobes/models/__init__.py +0 -0
  148. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/nnet/CNN.py +0 -0
  149. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/nnet/containers.py +0 -0
  150. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/nnet/linear.py +0 -0
  151. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/nnet/normalization.py +0 -0
  152. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/nnet/pooling.py +0 -0
  153. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/processing/__init__.py +0 -0
  154. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/processing/features.py +0 -0
  155. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/utils/__init__.py +0 -0
  156. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/utils/_workarounds.py +0 -0
  157. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/utils/autocast.py +0 -0
  158. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/utils/checkpoints.py +0 -0
  159. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/utils/fetching.py +0 -0
  160. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/utils/filter_analysis.py +0 -0
  161. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/utils/logger.py +0 -0
  162. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/utils/parameter_transfer.py +0 -0
  163. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/speechbrain/utils/run_opts.py +0 -0
  164. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/LICENSE +0 -0
  165. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/__init__.py +0 -0
  166. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/model/__init__.py +0 -0
  167. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/model/base.py +0 -0
  168. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/model/spin.py +0 -0
  169. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/nn/__init__.py +0 -0
  170. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/nn/dnn.py +0 -0
  171. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/nn/hubert.py +0 -0
  172. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/nn/swav_vq_dis.py +0 -0
  173. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/nn/wavlm.py +0 -0
  174. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/util/__init__.py +0 -0
  175. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/util/model_utils.py +0 -0
  176. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/spin/util/padding.py +0 -0
  177. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/timm/LICENSE +0 -0
  178. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/timm/__init__.py +0 -0
  179. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/timm/layers/__init__.py +0 -0
  180. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/timm/layers/drop.py +0 -0
  181. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/timm/layers/helpers.py +0 -0
  182. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/third_party/timm/layers/mlp.py +0 -0
  183. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/utils/__init__.py +0 -0
  184. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/utils/paths.py +0 -0
  185. {lfeats-0.1.4 → lfeats-0.2.0}/lfeats/utils/validation.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: lfeats
3
- Version: 0.1.4
3
+ Version: 0.2.0
4
4
  Summary: A unified interface to extract hidden representations from speech foundation models
5
5
  Project-URL: Documentation, https://takenori-y.github.io/lfeats/stable/
6
6
  Project-URL: Source, https://github.com/takenori-y/lfeats
@@ -18,6 +18,7 @@ Classifier: Programming Language :: Python :: 3.12
18
18
  Classifier: Programming Language :: Python :: 3.13
19
19
  Classifier: Programming Language :: Python :: 3.14
20
20
  Requires-Python: >=3.10
21
+ Requires-Dist: filelock>=3.10.0
21
22
  Requires-Dist: huggingface-hub>=0.23.0
22
23
  Requires-Dist: hydra-core>=1.3.0
23
24
  Requires-Dist: hyperpyyaml>=0.0.1
@@ -90,6 +90,18 @@ class Extractor:
90
90
  """
91
91
  self.model_manager.get_model().load(self.cache_dir, quiet)
92
92
 
93
+ def to(self, device: str) -> None:
94
+ """Move the model to the specified device.
95
+
96
+ Parameters
97
+ ----------
98
+ device : str
99
+ The device to move the model to (e.g., 'cpu' or 'cuda').
100
+
101
+ """
102
+ self.model_manager.to(device)
103
+ self.resampler_manager.to(device)
104
+
93
105
  def __call__(
94
106
  self,
95
107
  source: np.ndarray | torch.Tensor | Audio,
@@ -45,6 +45,17 @@ class Resampler:
45
45
  f"Supported resamplers are: {[k for k in RESAMPLER_MAP.keys()]}"
46
46
  ) from e
47
47
 
48
+ def to(self, device: str) -> None:
49
+ """Move the resampler to the specified device.
50
+
51
+ Parameters
52
+ ----------
53
+ device : str
54
+ The device to move the resampler to (e.g., 'cpu' or 'cuda').
55
+
56
+ """
57
+ self.resampler_manager.to(device)
58
+
48
59
  def __call__(
49
60
  self,
50
61
  source: np.ndarray | torch.Tensor | Audio,
@@ -25,6 +25,7 @@ class BaseModel(ABC):
25
25
  """
26
26
  self.device = device
27
27
 
28
+ self.model = None
28
29
  self._model_id = None # To be defined in subclasses
29
30
 
30
31
  @abstractmethod
@@ -42,6 +43,24 @@ class BaseModel(ABC):
42
43
  """
43
44
  raise NotImplementedError
44
45
 
46
+ def to(self, device: str) -> None:
47
+ """Move the model to the specified device.
48
+
49
+ Parameters
50
+ ----------
51
+ device : str
52
+ The device to move the model to (e.g., 'cpu' or 'cuda').
53
+
54
+ """
55
+ self.device = device
56
+ if self.model is not None and hasattr(self.model, "to"):
57
+ self.model.to(device)
58
+ if hasattr(self.model, "device"):
59
+ try:
60
+ self.model.device = device
61
+ except AttributeError:
62
+ pass
63
+
45
64
  def extract_features(self, audio: Audio, layers: list[int]) -> Features:
46
65
  """Extract features from the input audio data.
47
66
 
@@ -68,8 +68,6 @@ class ContentVecModel(FrameLevelFeatureModel):
68
68
  )
69
69
  self._model_id = f"contentvec-{self.variant.value}"
70
70
 
71
- self.model = None
72
-
73
71
  def load(self, model_dir: str, quiet: bool = False) -> None:
74
72
  """Load the model from the specified directory.
75
73
 
@@ -58,9 +58,6 @@ class Data2VecModel(FrameLevelFeatureModel):
58
58
  self.variant = validate_enum(variant, Data2VecVariant, Data2VecVariant.BASE)
59
59
  self._model_id = f"data2vec-{self.variant.value}"
60
60
 
61
- self.processor = None
62
- self.model = None
63
-
64
61
  def load(self, model_dir: str, quiet: bool = False) -> None:
65
62
  """Load the model from the specified directory.
66
63
 
@@ -58,9 +58,6 @@ class Data2Vec2Model(FrameLevelFeatureModel):
58
58
  self.variant = validate_enum(variant, Data2Vec2Variant, Data2Vec2Variant.BASE)
59
59
  self._model_id = f"data2vec2-{self.variant.value}"
60
60
 
61
- self.processor = None
62
- self.model = None
63
-
64
61
  def load(self, model_dir: str, quiet: bool = False) -> None:
65
62
  """Load the model from the specified directory.
66
63
 
@@ -40,8 +40,6 @@ class EcapaTDNNModel(UtteranceLevelFeatureModel):
40
40
  self.variant = validate_enum(variant, EcapaTDNNVariant, EcapaTDNNVariant.BASE)
41
41
  self._model_id = f"ecapa-tdnn-{self.variant.value}"
42
42
 
43
- self.model = None
44
-
45
43
  def load(self, model_dir: str, quiet: bool = False) -> None:
46
44
  """Load the model from the specified directory.
47
45
 
@@ -78,11 +76,11 @@ class EcapaTDNNModel(UtteranceLevelFeatureModel):
78
76
  self.model = EncoderClassifier.from_hparams(
79
77
  source="speechbrain/spkrec-ecapa-voxceleb",
80
78
  fetch_config=fetch_config,
79
+ run_opts={"device": self.device},
81
80
  )
82
81
  if self.model is None:
83
82
  raise RuntimeError("Failed to load the model.")
84
83
  self.model.eval()
85
- self.model.to(self.device)
86
84
 
87
85
  def extract_features_impl(self, audio: Audio, layers: list[int]) -> Features:
88
86
  """Extract features from the input audio using the model.
@@ -55,8 +55,6 @@ class Emotion2VecModel(FrameLevelFeatureModel):
55
55
  )
56
56
  self._model_id = f"emotion2vec-{self.variant.value}"
57
57
 
58
- self.model = None
59
-
60
58
  def load(self, model_dir: str, quiet: bool = False) -> None:
61
59
  """Load the model from the specified directory.
62
60
 
@@ -78,7 +76,7 @@ class Emotion2VecModel(FrameLevelFeatureModel):
78
76
  repo_id=repo_id,
79
77
  filename=filename,
80
78
  repo_type="model",
81
- local_dir=model_dir,
79
+ cache_dir=model_dir,
82
80
  )
83
81
 
84
82
  setup_third_party_path()
@@ -58,8 +58,6 @@ class Emotion2VecPlusModel(FrameLevelFeatureModel):
58
58
  )
59
59
  self._model_id = f"emotion2vec+-{self.variant.value}"
60
60
 
61
- self.model = None
62
-
63
61
  def load(self, model_dir: str, quiet: bool = False) -> None:
64
62
  """Load the model from the specified directory.
65
63
 
@@ -81,13 +79,13 @@ class Emotion2VecPlusModel(FrameLevelFeatureModel):
81
79
  repo_id=repo_id,
82
80
  filename="model.pt",
83
81
  repo_type="model",
84
- local_dir=os.path.join(model_dir, sanitize(self.model_id)),
82
+ cache_dir=os.path.join(model_dir, sanitize(self.model_id)),
85
83
  )
86
84
  config = hf_hub_download(
87
85
  repo_id=repo_id,
88
86
  filename="config.yaml",
89
87
  repo_type="model",
90
- local_dir=os.path.join(model_dir, sanitize(self.model_id)),
88
+ cache_dir=os.path.join(model_dir, sanitize(self.model_id)),
91
89
  )
92
90
 
93
91
  from hyperpyyaml import load_hyperpyyaml
@@ -60,8 +60,6 @@ class HuBERTModel(FrameLevelFeatureModel):
60
60
  self.variant = validate_enum(variant, HuBERTVariant, HuBERTVariant.BASE)
61
61
  self._model_id = f"hubert-{self.variant.value}"
62
62
 
63
- self.model = None
64
-
65
63
  def load(self, model_dir: str, quiet: bool = False) -> None:
66
64
  """Load the model from the specified directory.
67
65
 
@@ -32,6 +32,19 @@ class ModelManager:
32
32
 
33
33
  self._cache: dict[str, BaseModel] = {}
34
34
 
35
+ def to(self, device: str) -> None:
36
+ """Move all models to the specified device.
37
+
38
+ Parameters
39
+ ----------
40
+ device : str
41
+ The device to move the models to (e.g., 'cpu' or 'cuda').
42
+
43
+ """
44
+ self.device = device
45
+ for model in self._cache.values():
46
+ model.to(device)
47
+
35
48
  def get_model(self) -> BaseModel:
36
49
  """Get the model instance.
37
50
 
@@ -66,8 +66,6 @@ class NeXtTDNNModel(UtteranceLevelFeatureModel):
66
66
  self.variant = validate_enum(variant, NeXtTDNNVariant, NeXtTDNNVariant.BASE_V2)
67
67
  self._model_id = f"next-tdnn-{self.variant.value}"
68
68
 
69
- self.model = None
70
-
71
69
  def load(self, model_dir: str, quiet: bool = False) -> None:
72
70
  """Load the model from the specified directory.
73
71
 
@@ -63,8 +63,6 @@ class RSpinModel(FrameLevelFeatureModel):
63
63
  self.variant = validate_enum(variant, RSpinVariant, RSpinVariant.WAVLM_256)
64
64
  self._model_id = f"r-spin-{self.variant.value}"
65
65
 
66
- self.model = None
67
-
68
66
  def load(self, model_dir: str, quiet: bool = False) -> None:
69
67
  """Load the model from the specified directory.
70
68
 
@@ -40,8 +40,6 @@ class RVectorModel(UtteranceLevelFeatureModel):
40
40
  self.variant = validate_enum(variant, RVectorVariant, RVectorVariant.BASE)
41
41
  self._model_id = f"r-vector-{self.variant.value}"
42
42
 
43
- self.model = None
44
-
45
43
  def load(self, model_dir: str, quiet: bool = False) -> None:
46
44
  """Load the model from the specified directory.
47
45
 
@@ -78,11 +76,11 @@ class RVectorModel(UtteranceLevelFeatureModel):
78
76
  self.model = EncoderClassifier.from_hparams(
79
77
  source="speechbrain/spkrec-resnet-voxceleb",
80
78
  fetch_config=fetch_config,
79
+ run_opts={"device": self.device},
81
80
  )
82
81
  if self.model is None:
83
82
  raise RuntimeError("Failed to load the model.")
84
83
  self.model.eval()
85
- self.model.to(self.device)
86
84
 
87
85
  def extract_features_impl(self, audio: Audio, layers: list[int]) -> Features:
88
86
  """Extract features from the input audio using the model.
@@ -4,12 +4,11 @@
4
4
  """A module for the SpidR model."""
5
5
 
6
6
  from enum import Enum
7
- from typing import Any
8
7
 
9
8
  import torch
10
9
 
11
10
  from ..interfaces.types import Audio, Features
12
- from ..utils.io import set_torch_hub_dir
11
+ from ..utils.io import safe_torch_hub_load
13
12
  from ..utils.validation import validate_enum
14
13
  from .base import FrameLevelFeatureModel
15
14
 
@@ -40,8 +39,6 @@ class SpidRModel(FrameLevelFeatureModel):
40
39
  self.variant = validate_enum(variant, SpidRVariant, SpidRVariant.BASE)
41
40
  self._model_id = f"spidr-{self.variant.value}"
42
41
 
43
- self.model = None
44
-
45
42
  def load(self, model_dir: str, quiet: bool = False) -> None:
46
43
  """Load the model from the specified directory.
47
44
 
@@ -57,10 +54,9 @@ class SpidRModel(FrameLevelFeatureModel):
57
54
  if self.model is not None:
58
55
  return
59
56
 
60
- with set_torch_hub_dir(model_dir):
61
- self.model: Any = torch.hub.load(
62
- "facebookresearch/spidr", "spidr_base", verbose=not quiet
63
- )
57
+ self.model = safe_torch_hub_load(
58
+ "facebookresearch/spidr", "spidr_base", model_dir, quiet=quiet
59
+ )
64
60
  self.model.eval()
65
61
  self.model.to(self.device)
66
62
 
@@ -117,8 +117,6 @@ class SpinModel(FrameLevelFeatureModel):
117
117
  self.variant = validate_enum(variant, SpinVariant, SpinVariant.HUBERT_256)
118
118
  self._model_id = f"spin-{self.variant.value}"
119
119
 
120
- self.model = None
121
-
122
120
  def load(self, model_dir: str, quiet: bool = False) -> None:
123
121
  """Load the model from the specified directory.
124
122
 
@@ -139,7 +137,7 @@ class SpinModel(FrameLevelFeatureModel):
139
137
  repo_id="vectominist/spin_ckpt",
140
138
  filename=self.variant.checkpoint_filename,
141
139
  repo_type="dataset",
142
- local_dir=model_dir,
140
+ cache_dir=model_dir,
143
141
  )
144
142
 
145
143
  with lightning_mock_context():
@@ -147,11 +145,12 @@ class SpinModel(FrameLevelFeatureModel):
147
145
  model_path, map_location=torch.device("cpu"), weights_only=False
148
146
  )
149
147
 
150
- from lfeats.third_party.s3prl.util.download import set_dir
148
+ from lfeats.third_party.s3prl.util.download import set_dir, set_progress
151
149
  from lfeats.third_party.spin.model import SpinModel as _SpinModel
152
150
  from lfeats.third_party.spin.util import len_to_padding
153
151
 
154
152
  set_dir(model_dir)
153
+ set_progress(not quiet)
155
154
  self.model = _SpinModel(checkpoint["hyper_parameters"])
156
155
  self.model.load_state_dict(checkpoint["state_dict"])
157
156
  self.model.eval()
@@ -63,6 +63,19 @@ class SSLZipModel(FrameLevelFeatureModel):
63
63
  self.upstream = HuBERTModel(variant=HuBERTVariant.BASE.value, device=device)
64
64
  self.model = None
65
65
 
66
+ def to(self, device: str) -> None:
67
+ """Move the model to the specified device.
68
+
69
+ Parameters
70
+ ----------
71
+ device : str
72
+ The device to move the model to (e.g., 'cpu' or 'cuda'). Note that the ONNX
73
+ model will be loaded on the specified device when calling the 'load' method.
74
+
75
+ """
76
+ self.device = device
77
+ self.upstream.to(device)
78
+
66
79
  def load(self, model_dir: str, quiet: bool = False) -> None:
67
80
  """Load the model from the specified directory.
68
81
 
@@ -86,7 +99,7 @@ class SSLZipModel(FrameLevelFeatureModel):
86
99
  repo_id=repo_id,
87
100
  filename=filename,
88
101
  repo_type="model",
89
- local_dir=model_dir,
102
+ cache_dir=model_dir,
90
103
  )
91
104
 
92
105
  import onnxruntime as ort
@@ -56,8 +56,6 @@ class UniSpeechSATModel(FrameLevelFeatureModel):
56
56
  )
57
57
  self._model_id = f"unispeech-sat-{self.variant.value}"
58
58
 
59
- self.model = None
60
-
61
59
  def load(self, model_dir: str, quiet: bool = False) -> None:
62
60
  """Load the model from the specified directory.
63
61
 
@@ -64,9 +64,6 @@ class Wav2Vec2Model(FrameLevelFeatureModel):
64
64
  self.variant = validate_enum(variant, Wav2Vec2Variant, Wav2Vec2Variant.BASE)
65
65
  self._model_id = f"wav2vec2-{self.variant.value}"
66
66
 
67
- self.processor = None
68
- self.model = None
69
-
70
67
  def load(self, model_dir: str, quiet: bool = False) -> None:
71
68
  """Load the model from the specified directory.
72
69
 
@@ -54,8 +54,6 @@ class WavLMModel(FrameLevelFeatureModel):
54
54
  self.variant = validate_enum(variant, WavLMVariant, WavLMVariant.BASE_PLUS)
55
55
  self._model_id = f"wavlm-{self.variant.value}"
56
56
 
57
- self.model = None
58
-
59
57
  def load(self, model_dir: str, quiet: bool = False) -> None:
60
58
  """Load the model from the specified directory.
61
59
 
@@ -58,7 +58,6 @@ class WhisperModel(FrameLevelFeatureModel):
58
58
  self._model_id = f"whisper-{self.variant.value}"
59
59
 
60
60
  self.processor = None
61
- self.model = None
62
61
 
63
62
  def load(self, model_dir: str, quiet: bool = False) -> None:
64
63
  """Load the model from the specified directory.
@@ -40,8 +40,6 @@ class XVectorModel(UtteranceLevelFeatureModel):
40
40
  self.variant = validate_enum(variant, XVectorVariant, XVectorVariant.BASE)
41
41
  self._model_id = f"x-vector-{self.variant.value}"
42
42
 
43
- self.model = None
44
-
45
43
  def load(self, model_dir: str, quiet: bool = False) -> None:
46
44
  """Load the model from the specified directory.
47
45
 
@@ -78,11 +76,11 @@ class XVectorModel(UtteranceLevelFeatureModel):
78
76
  self.model = EncoderClassifier.from_hparams(
79
77
  source="speechbrain/spkrec-xvect-voxceleb",
80
78
  fetch_config=fetch_config,
79
+ run_opts={"device": self.device},
81
80
  )
82
81
  if self.model is None:
83
82
  raise RuntimeError("Failed to load the model.")
84
83
  self.model.eval()
85
- self.model.to(self.device)
86
84
 
87
85
  def extract_features_impl(self, audio: Audio, layers: list[int]) -> Features:
88
86
  """Extract features from the input audio using the model.
@@ -35,6 +35,17 @@ class BaseResampler(ABC):
35
35
  self.dst_rate = dst_rate
36
36
  self.device = device
37
37
 
38
+ def to(self, device: str) -> None:
39
+ """Move the resampler to the specified device.
40
+
41
+ Parameters
42
+ ----------
43
+ device : str
44
+ The device to move the resampler to (e.g., 'cpu' or 'cuda').
45
+
46
+ """
47
+ self.device = device
48
+
38
49
  def resample(self, audio: Audio) -> Audio:
39
50
  """Resample the given audio to the target sample rate.
40
51
 
@@ -56,7 +56,19 @@ class LilFilterResampler(BaseResampler):
56
56
  dst_rate,
57
57
  dtype=torch.float32,
58
58
  )
59
- self.resampler.weights = self.resampler.weights.to(self.device)
59
+ self.to(device)
60
+
61
+ def to(self, device: str) -> None:
62
+ """Move the resampler to the specified device.
63
+
64
+ Parameters
65
+ ----------
66
+ device : str
67
+ The device to move the resampler to (e.g., 'cpu' or 'cuda').
68
+
69
+ """
70
+ super().to(device)
71
+ self.resampler.weights = self.resampler.weights.to(device)
60
72
 
61
73
  def resample_impl(self, audio: Audio) -> Audio:
62
74
  """Resample the given audio to the target sample rate.
@@ -32,6 +32,19 @@ class ResamplerManager:
32
32
 
33
33
  self._cache: dict[tuple[int, int], BaseResampler] = {}
34
34
 
35
+ def to(self, device: str) -> None:
36
+ """Move all resamplers to the specified device.
37
+
38
+ Parameters
39
+ ----------
40
+ device : str
41
+ The device to move the resamplers to (e.g., 'cpu' or 'cuda').
42
+
43
+ """
44
+ self.device = device
45
+ for resampler in self._cache.values():
46
+ resampler.to(device)
47
+
35
48
  def get_resampler(self, src_rate: int, dst_rate: int) -> BaseResampler:
36
49
  """Get the resampler instance.
37
50
 
@@ -73,7 +73,20 @@ class TorchAudioResampler(BaseResampler):
73
73
 
74
74
  self.resampler = torchaudio.transforms.Resample(
75
75
  orig_freq=src_rate, new_freq=dst_rate, dtype=torch.float32, **params
76
- ).to(device)
76
+ )
77
+ self.to(device)
78
+
79
+ def to(self, device: str) -> None:
80
+ """Move the resampler to the specified device.
81
+
82
+ Parameters
83
+ ----------
84
+ device : str
85
+ The device to move the resampler to (e.g., 'cpu' or 'cuda').
86
+
87
+ """
88
+ super().to(device)
89
+ self.resampler.to(device)
77
90
 
78
91
  def resample_impl(self, audio: Audio) -> Audio:
79
92
  """Resample the given audio to the target sample rate.
@@ -24,10 +24,13 @@ logger = logging.getLogger(__name__)
24
24
 
25
25
 
26
26
  _download_dir = Path.home() / ".cache" / "s3prl" / "download"
27
+ _progress = True
27
28
 
28
29
  __all__ = [
29
30
  "get_dir",
30
31
  "set_dir",
32
+ "get_progress",
33
+ "set_progress",
31
34
  "download",
32
35
  "urls_to_filepaths",
33
36
  ]
@@ -43,6 +46,15 @@ def set_dir(d):
43
46
  _download_dir = Path(d)
44
47
 
45
48
 
49
+ def get_progress():
50
+ return _progress
51
+
52
+
53
+ def set_progress(progress):
54
+ global _progress
55
+ _progress = progress
56
+
57
+
46
58
  def _download_url_to_file(url, dst, hash_prefix=None, progress=True):
47
59
  """
48
60
  This function is not thread-safe. Please ensure only a single
@@ -69,8 +81,9 @@ def _download_url_to_file(url, dst, hash_prefix=None, progress=True):
69
81
  if hash_prefix is not None:
70
82
  sha256 = hashlib.sha256()
71
83
 
72
- tqdm.write(f"Downloading: {url}", file=sys.stderr)
73
- tqdm.write(f"Destination: {dst}", file=sys.stderr)
84
+ if progress:
85
+ tqdm.write(f"Downloading: {url}", file=sys.stderr)
86
+ tqdm.write(f"Destination: {dst}", file=sys.stderr)
74
87
  with tqdm(
75
88
  total=file_size,
76
89
  disable=not progress,
@@ -120,12 +133,13 @@ def _download_url_to_file_requests(url, dst, hash_prefix=None, progress=True):
120
133
  if hash_prefix is not None:
121
134
  sha256 = hashlib.sha256()
122
135
 
123
- tqdm.write(
124
- f"urllib.Request method failed. Trying using another method...",
125
- file=sys.stderr,
126
- )
127
- tqdm.write(f"Downloading: {url}", file=sys.stderr)
128
- tqdm.write(f"Destination: {dst}", file=sys.stderr)
136
+ if progress:
137
+ tqdm.write(
138
+ f"urllib.Request method failed. Trying using another method...",
139
+ file=sys.stderr,
140
+ )
141
+ tqdm.write(f"Downloading: {url}", file=sys.stderr)
142
+ tqdm.write(f"Destination: {dst}", file=sys.stderr)
129
143
  with tqdm(
130
144
  total=file_size,
131
145
  disable=not progress,
@@ -176,9 +190,9 @@ def _download(filepath: Path, url, refresh: bool, new_enough_secs: float = 2.0):
176
190
  refresh and (time.time() - os.path.getmtime(filepath)) > new_enough_secs
177
191
  ):
178
192
  try:
179
- _download_url_to_file(url, filepath)
193
+ _download_url_to_file(url, filepath, progress=_progress)
180
194
  except:
181
- _download_url_to_file_requests(url, filepath)
195
+ _download_url_to_file_requests(url, filepath, progress=_progress)
182
196
 
183
197
  logger.info(f"Using URL's local file: {filepath}")
184
198
 
@@ -8,7 +8,7 @@ Authors:
8
8
  """
9
9
 
10
10
  # import datetime
11
- # import os
11
+ import os
12
12
  from functools import wraps
13
13
  from typing import Optional
14
14
 
@@ -4,13 +4,16 @@
4
4
  """I/O utilities."""
5
5
 
6
6
  import logging
7
+ import os
7
8
  from collections.abc import Generator
8
9
  from contextlib import contextmanager
9
10
  from importlib.metadata import PackageNotFoundError, version
11
+ from typing import Any
10
12
 
11
13
  import soundfile as sf
12
14
  import torch
13
15
  import torchaudio
16
+ from filelock import FileLock
14
17
 
15
18
  HF_HTTP_LOGGER = "huggingface_hub.utils._http"
16
19
 
@@ -108,14 +111,50 @@ def download_file(url: str, download_dir: str, quiet: bool = False) -> str:
108
111
  """
109
112
  from parfive import Downloader
110
113
 
111
- dl = Downloader(progress=not quiet)
114
+ basename = os.path.basename(url)
115
+ lock_path = os.path.join(download_dir, f"{basename}.lock")
112
116
 
113
- dl.enqueue_file(url, path=download_dir)
117
+ with FileLock(lock_path):
118
+ dl = Downloader(progress=not quiet)
114
119
 
115
- results = dl.download()
116
- if len(results) == 0:
117
- raise RuntimeError(f"Failed to download file: {results.errors}")
118
- return results[0]
120
+ dl.enqueue_file(url, path=download_dir)
121
+
122
+ results = dl.download()
123
+ if len(results) == 0:
124
+ raise RuntimeError(f"Failed to download file: {results.errors}")
125
+ return results[0]
126
+
127
+
128
+ def safe_torch_hub_load(
129
+ repo_or_dir: str, model: str, download_dir: str, quiet: bool = False
130
+ ) -> Any:
131
+ """Safely load a model from PyTorch Hub.
132
+
133
+ Parameters
134
+ ----------
135
+ repo_or_dir : str
136
+ The repository or directory to load the model from.
137
+
138
+ model : str
139
+ The name of the model defined in the repository.
140
+
141
+ download_dir : str
142
+ The directory to use for caching the downloaded model files.
143
+
144
+ quiet : bool, optional
145
+ Whether to suppress output during the loading process.
146
+
147
+ Returns
148
+ -------
149
+ out : Any
150
+ The loaded object.
151
+
152
+ """
153
+ lock_path = os.path.join(download_dir, f"{model}.lock")
154
+
155
+ with FileLock(lock_path):
156
+ with set_torch_hub_dir(download_dir):
157
+ return torch.hub.load(repo_or_dir, model, verbose=not quiet)
119
158
 
120
159
 
121
160
  def load_audio(path: str) -> tuple[torch.Tensor, int]: