Class: OnnxRuby::Reranker

Inherits:
Object
  • Object
show all
Includes:
TokenizerSupport
Defined in:
lib/onnx_ruby/reranker.rb

Instance Attribute Summary collapse

Instance Method Summary collapse

Constructor Details

#initialize(model_path, tokenizer: nil, **session_opts) ⇒ Reranker

Returns a new instance of Reranker.



9
10
11
12
# File 'lib/onnx_ruby/reranker.rb', line 9

def initialize(model_path, tokenizer: nil, **session_opts)
  @session = Session.new(model_path, **session_opts)
  @tokenizer = resolve_tokenizer(tokenizer)
end

Instance Attribute Details

#sessionObject (readonly)

Returns the value of attribute session.



7
8
9
# File 'lib/onnx_ruby/reranker.rb', line 7

def session
  @session
end

Instance Method Details

#rerank(query, documents) ⇒ Array<Hash>

Rerank documents by relevance to a query

Parameters:

  • query (String)

    the query text (requires tokenizer)

  • documents (Array<String>)

    documents to rerank

Returns:

  • (Array<Hash>)

    sorted array of { document:, score:, index: }

Raises:



18
19
20
21
22
23
24
25
26
27
# File 'lib/onnx_ruby/reranker.rb', line 18

def rerank(query, documents)
  raise Error, "tokenizer is required for reranking" unless @tokenizer

  pairs = documents.map { |doc| [query, doc] }
  scores = score_pairs(pairs)

  documents.each_with_index.map do |doc, i|
    { document: doc, score: scores[i], index: i }
  end.sort_by { |r| -r[:score] }
end

#score(input_ids:, attention_mask:) ⇒ Array<Float>

Score query-document pairs with pre-tokenized inputs

Parameters:

  • input_ids (Array<Array<Integer>>)

    batch of token ID sequences

  • attention_mask (Array<Array<Integer>>)

    batch of attention masks

Returns:

  • (Array<Float>)

    relevance scores



33
34
35
36
37
38
# File 'lib/onnx_ruby/reranker.rb', line 33

def score(input_ids:, attention_mask:)
  feed = build_feed(input_ids, attention_mask)
  result = @session.run(feed)
  raw_scores = find_output(result, %w[scores logits output])
  raw_scores.map { |row| row.is_a?(Array) ? row.first : row }
end