import math
from typing import Dict, Tuple
class TransformerFLOPsEstimator:
"""Compute FLOPs for Transformer-based LLM rerankers."""
def __init__(self, model_config: Dict):
self.vocab_size = model_config.get("vocab_size", 50000)
self.hidden_dim = model_config.get("hidden_dim", 768)
self.num_layers = model_config.get("num_layers", 12)
self.num_heads = model_config.get("num_heads", 12)
self.ffn_dim = model_config.get("ffn_dim", 3072)
self.use_gqa = model_config.get("grouped_query_attention", False)
self.num_kv_heads = model_config.get("num_kv_heads", 1) if self.use_gqa else self.num_heads
def estimate_single_pass(self, seq_len: int, batch_size: int = 1) -> float:
"""
Estimate FLOPs for single forward pass through model.
Args:
seq_len: sequence length (input + output tokens)
batch_size: batch size
Returns:
estimated FLOPs
"""
total_flops = 0
embedding_flops = batch_size * seq_len * self.vocab_size * self.hidden_dim
total_flops += embedding_flops
for _ in range(self.num_layers):
q_proj_flops = batch_size * seq_len * self.hidden_dim * self.hidden_dim
total_flops += 2 * q_proj_flops
kv_proj_flops = batch_size * seq_len * self.hidden_dim * self.hidden_dim
total_flops += 2 * 2 * kv_proj_flops
head_dim = self.hidden_dim // self.num_heads
attention_flops = batch_size * self.num_heads * seq_len * seq_len * head_dim
total_flops += 2 * attention_flops
attn_output_flops = batch_size * self.num_heads * seq_len * seq_len * head_dim
total_flops += 2 * attn_output_flops
output_proj_flops = batch_size * seq_len * self.hidden_dim * self.hidden_dim
total_flops += 2 * output_proj_flops
ffn1_flops = batch_size * seq_len * self.hidden_dim * self.ffn_dim
total_flops += 2 * ffn1_flops
ffn2_flops = batch_size * seq_len * self.ffn_dim * self.hidden_dim
total_flops += 2 * ffn2_flops
output_vocab_flops = batch_size * seq_len * self.hidden_dim * self.vocab_size
total_flops += 2 * output_vocab_flops
return total_flops
def estimate_reranking_pass(
self,
query_len: int,
document_len: int,
num_documents: int,
batch_size: int = 1,
reranking_mode: str = "pointwise"
) -> float:
"""
Estimate FLOPs for reranking operation.
Args:
query_len: query token length
document_len: document token length
num_documents: number of documents to rank
batch_size: batch size
reranking_mode: "pointwise", "listwise", or "pairwise"
Returns:
estimated FLOPs
"""
if reranking_mode == "pointwise":
total_seq_len = query_len + document_len
total_pairs = num_documents
total_flops = 0
for _ in range(total_pairs):
total_flops += self.estimate_single_pass(total_seq_len, batch_size)
return total_flops
elif reranking_mode == "listwise":
total_seq_len = query_len + (document_len * num_documents)
return self.estimate_single_pass(total_seq_len, batch_size)
elif reranking_mode == "pairwise":
total_seq_len = query_len + document_len
num_pairs = (num_documents * (num_documents - 1)) // 2
return num_pairs * self.estimate_single_pass(total_seq_len, batch_size)
else:
raise ValueError(f"Unknown reranking mode: {reranking_mode}")
def estimate_with_gqa(self, seq_len: int, batch_size: int = 1) -> float:
"""
Estimate FLOPs with grouped-query attention (more efficient than MHA).
Reduces KV computation by factor of (num_heads / num_kv_heads).
"""
gqa_factor = self.num_heads / self.num_kv_heads
standard_flops = self.estimate_single_pass(seq_len, batch_size)
reduction = standard_flops * (1 - 1/gqa_factor) * 0.3
return standard_flops - reduction
import numpy as np
from typing import List, Tuple
class RankingQualityMetrics:
"""Compute ranking metrics per unit of computation."""
@staticmethod
def compute_mrr(rankings: List[List[int]], true_relevant: List[int]) -> float:
"""
Mean Reciprocal Rank: 1/rank of first relevant item.
Args:
rankings: list of ranked document indices per query
true_relevant: relevant document indices per query
Returns:
MRR score
"""
mrr_scores = []
for rank, doc_idx in enumerate(rankings, 1):
if doc_idx in true_relevant:
mrr_scores.append(1.0 / rank)
break
return np.mean(mrr_scores) if mrr_scores else 0.0
@staticmethod
def compute_ndcg(rankings: List[List[int]], true_relevant: List[int], k: int = 10) -> float:
"""
Normalized Discounted Cumulative Gain.
Args:
rankings: ranked document indices
true_relevant: relevant document indices
k: truncate at rank k
Returns:
NDCG@k score
"""
dcg = 0.0
for rank, doc_idx in enumerate(rankings[:k], 1):
if doc_idx in true_relevant:
dcg += 1.0 / math.log2(rank + 1)
ideal_dcg = sum(1.0 / math.log2(i + 1) for i in range(min(len(true_relevant), k)))
return dcg / ideal_dcg if ideal_dcg > 0 else 0.0
@staticmethod
def compute_hits_at_k(rankings: List[int], true_relevant: List[int], k: int = 10) -> float:
"""Hits@k: fraction of queries with at least one relevant document in top-k."""
return float(any(doc in true_relevant for doc in rankings[:k]))
class EfficiencyEffectivenessEvaluator:
def __init__(self, flops_estimator: TransformerFLOPsEstimator):
self.flops_estimator = flops_estimator
self.metrics = RankingQualityMetrics()
def evaluate_reranker(
self,
query_len: int,
document_len: int,
num_documents: int,
rankings: List[List[int]],
true_relevant: List[List[int]],
reranking_mode: str = "pointwise"
) -> Dict[str, float]:
"""
Compute efficiency-effectiveness metrics for reranker.
Returns metrics per PetaFLOP (10^15 FLOPs).
"""
flops = self.flops_estimator.estimate_reranking_pass(
query_len, document_len, num_documents, batch_size=1,
reranking_mode=reranking_mode
)
petaflops = flops / 1e15
mrr_scores = []
ndcg_scores = []
hits_at_10 = []
for i, ranking in enumerate(rankings):
relevant = true_relevant[i]
mrr = self.metrics.compute_mrr([ranking], relevant)
ndcg = self.metrics.compute_ndcg(ranking, relevant, k=10)
hits = self.metrics.compute_hits_at_k(ranking, relevant, k=10)
mrr_scores.append(mrr)
ndcg_scores.append(ndcg)
hits_at_10.append(hits)
avg_mrr = np.mean(mrr_scores)
avg_ndcg = np.mean(ndcg_scores)
avg_hits = np.mean(hits_at_10)
return {
"flops": flops,
"petaflops": petaflops,
"mrr": avg_mrr,
"ndcg@10": avg_ndcg,
"hits@10": avg_hits,
"rpp_mrr": avg_mrr / petaflops,
"rpp_ndcg": avg_ndcg / petaflops,
"queries_per_petaflop": 1.0 / petaflops,
}
def compare_reranking_strategies(
self,
query_len: int,
document_len: int,
num_documents: int,
rankings_dict: Dict[str, List[List[int]]],
true_relevant: List[List[int]]
) -> Dict[str, Dict]:
"""Compare multiple reranking strategies (pointwise, listwise, pairwise)."""
results = {}
for strategy in ["pointwise", "listwise", "pairwise"]:
results[strategy] = self.evaluate_reranker(
query_len, document_len, num_documents,
rankings_dict[strategy], true_relevant,
reranking_mode=strategy
)
return results
class RerankerBenchmark:
def __init__(self):
self.results = []
def benchmark_model_variants(
self,
base_config: Dict,
model_sizes: List[int],
reranking_modes: List[str],
query_len: int = 10,
document_len: int = 100,
num_documents: int = 100
) -> Dict:
"""Benchmark different model sizes and reranking strategies."""
benchmark_results = {}
for size in model_sizes:
config = base_config.copy()
config["hidden_dim"] = size
config["ffn_dim"] = size * 4
estimator = TransformerFLOPsEstimator(config)
evaluator = EfficiencyEffectivenessEvaluator(estimator)
for mode in reranking_modes:
key = f"model_{size}_mode_{mode}"
flops = estimator.estimate_reranking_pass(
query_len, document_len, num_documents,
reranking_mode=mode
)
rankings = [[i for i in range(num_documents)]]
true_relevant = [[0, 1, 2]]
metrics = evaluator.evaluate_reranker(
query_len, document_len, num_documents,
rankings, true_relevant, mode
)
benchmark_results[key] = metrics
return benchmark_results
def summarize_efficiency_dominance(self, results: Dict) -> str:
"""Analyze which strategies are most efficient."""
summary = []
strategy_results = {}
for key, metrics in results.items():
strategy = key.split("_mode_")[1]
if strategy not in strategy_results:
strategy_results[strategy] = []
strategy_results[strategy].append((key, metrics))
for strategy, items in strategy_results.items():
avg_rpp = np.mean([m["rpp_mrr"] for _, m in items])
summary.append(f"{strategy}: avg RPP = {avg_rpp:.4f}")
return "\n".join(sorted(summary, reverse=True))