Class: RerankerRuby::ModelDownloader

Inherits:
Object
  • Object
show all
Defined in:
lib/reranker_ruby/model_downloader.rb

Constant Summary collapse

HF_BASE_URL =
"https://huggingface.co"
DEFAULT_CACHE_DIR =
File.join(Dir.home, ".cache", "reranker-ruby", "models")
ONNX_PATHS =

Known ONNX model paths for popular cross-encoder models

{
  "cross-encoder/ms-marco-MiniLM-L-6-v2" => "onnx/model.onnx",
  "cross-encoder/ms-marco-MiniLM-L-12-v2" => "onnx/model.onnx",
  "BAAI/bge-reranker-base" => "onnx/model.onnx",
  "BAAI/bge-reranker-large" => "onnx/model.onnx",
  "BAAI/bge-reranker-v2-m3" => "onnx/model.onnx"
}.freeze

Instance Method Summary collapse

Constructor Details

#initialize(cache_dir: DEFAULT_CACHE_DIR, token: nil) ⇒ ModelDownloader

Returns a new instance of ModelDownloader.



21
22
23
24
# File 'lib/reranker_ruby/model_downloader.rb', line 21

def initialize(cache_dir: DEFAULT_CACHE_DIR, token: nil)
  @cache_dir = cache_dir
  @token = token
end

Instance Method Details

#download(repo_id) ⇒ Hash

Downloads model and tokenizer files, returns paths

Returns:

  • (Hash) —

    { model_path:, tokenizer_path: }



28
29
30
31
32
33
34
35
36
37
38
39
40
41
# File 'lib/reranker_ruby/model_downloader.rb', line 28

def download(repo_id)
  model_dir = File.join(@cache_dir, repo_id.gsub("/", "--"))
  FileUtils.mkdir_p(model_dir)

  onnx_path = ONNX_PATHS.fetch(repo_id, "onnx/model.onnx")

  model_path = File.join(model_dir, "model.onnx")
  tokenizer_path = File.join(model_dir, "tokenizer.json")

  download_file(repo_id, onnx_path, model_path) unless File.exist?(model_path)
  download_file(repo_id, "tokenizer.json", tokenizer_path) unless File.exist?(tokenizer_path)

  { model_path: model_path, tokenizer_path: tokenizer_path }
end