Views
No views yet
Qwen3-Embedding-0.6B.1
2import argparse
3import torch
4import pandas as pd
5import pyarrow.parquet as pq
6from transformers import AutoTokenizer, AutoModelForSequenceClassification
7
8def percentile_ranks(scores):
9 # higher=better -> percentile in [0,1], 1.0 is best
10 order = torch.argsort(scores, descending=True)
11 ranks = torch.empty_like(order, dtype=torch.float)
12 ranks[order] = torch.arange(len(scores), dtype=torch.float)
13 denom = max(1, len(scores) - 1)
14 return 1.0 - ranks / denom
15
16@torch.no_grad()
17def batched_logits(texts, tokenizer, model, batch_size=64, max_length=512, device="cuda" if torch.cuda.is_available() else "cpu"):
18 model.to(device).eval()
19 out_scores = []
20 for i in range(0, len(texts), batch_size):
21 batch = texts[i:i+batch_size]
22 enc = tokenizer(batch, padding=True, truncation=True,
23 max_length=max_length, return_tensors="pt").to(device)
24 logits = model(**enc).logits.squeeze(-1) # (B,) for class_num=1
25 out_scores.append(logits.cpu())
26 return torch.cat(out_scores, dim=0)
27
28def main():
29 ap = argparse.ArgumentParser()
30 ap.add_argument("--parquet", required=True, help="Input parquet path with column 'text'.")
31 ap.add_argument("--rater1", default="pararater_rater1_en-ar", help="Rater1.")
32 ap.add_argument("--rater2", default="pararater_rater2_en-ar", help="Rater2.")
33 ap.add_argument("--save_parquet", default=None, help="Optional output parquet for kept samples.")
34 ap.add_argument("--batch_size", type=int, default=64)
35 ap.add_argument("--max_length", type=int, default=512)
36 args = ap.parse_args()
37
38 # 1) Load data
39 df = pq.read_table(args.parquet).to_pandas()
40 assert "text" in df.columns, "Parquet must have a 'text' column."
41 texts = df["text"].astype(str).tolist()
42
43 # 2) Load raters
44 tok = AutoTokenizer.from_pretrained(args.rater1, trust_remote_code=True)
45 r1 = AutoModelForSequenceClassification.from_pretrained(args.rater1, trust_remote_code=True)
46 r2 = AutoModelForSequenceClassification.from_pretrained(args.rater2, trust_remote_code=True)
47
48 # 3) Score -> percentile ranks
49 s1 = batched_logits(texts, tok, r1, batch_size=args.batch_size, max_length=args.max_length)
50 s2 = batched_logits(texts, tok, r2, batch_size=args.batch_size, max_length=args.max_length)
51 p1 = percentile_ranks(s1) # 1.0 best
52 p2 = percentile_ranks(s2)
53
54 # 4) Rule: keep if (p1 >= 0.6) and (p2 <= p1 - 0.2)
55 top = p1 >= 0.6
56 drop = p2 <= (p1 - 0.2)
57 keep_mask = (top & drop).numpy()
58
59 kept = df.loc[keep_mask]
60 print(f"Total: {len(df)} | Rater1 top-0.6: {int(top.sum().item())} | Kept(final): {keep_mask.sum()}")
61
62 if args.save_parquet:
63 kept.to_parquet(args.save_parquet, index=False)
64 print(f"Saved: {args.save_parquet}")
65
66if __name__ == "__main__":
67 main()
68