Shuu12121/owl_code_search_hard_negative_datasets-Pre_kd
Owl Code Search Hard Negative Datasets Knowledge Distillation (KD) ベースのハードネガティブ付きコード検索データセットです。コード検索モデルShuu12121/CodeSearch-ModernBERT-Crow-v3-large-len1024-Plusを教師モデルとして、各コメントと説明コメントのペアのデータセットから各クエリに対する関数の類似度スコアを計算し、ハードネガティブ(正解に類似しているが不正解の文書)を付与しています。 概要 目的: コード検索モデルの Contrastive Learning / Knowledge Distillation ファインチューニング 言語: Go, Java, JavaScript, PHP, Python, Ruby, Rust, TypeScript(8言語) 総サンプル数: 4,787,740 データサイズ: 8.73 GB(展開後) / 3.37 GB(ダウンロード時) フォーマット:… See the full description on the dataset page: https://huggingface.co/datasets/Shuu12121/owl_code_search_hard_negative_datasets-Pre_kd.
Owl Code Search Hard Negative Datasets
Knowledge Distillation (KD) ベースのハードネガティブ付きコード検索データセットです。 コード検索モデルShuu12121/CodeSearch-ModernBERT-Crow-v3-large-len1024-Plusを教師モデルとして、各コメントと説明コメントのペアのデータセットから各クエリに対する関数の類似度スコアを計算し、ハードネガティブ(正解に類似しているが不正解の文書)を付与しています。
概要
- 目的: コード検索モデルの Contrastive Learning / Knowledge Distillation ファインチューニング
- 言語: Go, Java, JavaScript, PHP, Python, Ruby, Rust, TypeScript(8言語)
- 総サンプル数: 4,787,740
- データサイズ: 8.73 GB(展開後) / 3.37 GB(ダウンロード時)
- フォーマット: Per-language config 形式(
scores_{lang},queries_{lang},documents_{lang})
データ構造
各言語ごとに 3 つの config が存在します:
queries_{lang}
各クエリ(自然言語による検索文)を格納。
documents_{lang}
各文書(ソースコード)を格納。
scores_{lang}
教師モデルによる類似度スコアを格納。各クエリに対して、スコア順にソートされた文書 ID リストとスコアリストを保持。
スコアの解釈: -scores[0]/document_ids[0]が正例(実際のペアだったもの) -score[0] = -1は正解が上位32件に検索結果が含まれていなかった場合
言語別統計
注意点
全データをメモリに載せようとするとOOMになる可能性があります!!
使い方
基本的な読み込み
from datasets import load_dataset
# Python の scores を読み込む
scores = load_dataset(
"Shuu12121/owl_code_search_hard_negative_datasets-Pre_kd",
name="scores_python",
split="train",
)
# Python の queries を読み込む
queries = load_dataset(
"Shuu12121/owl_code_search_hard_negative_datasets-Pre_kd",
name="queries_python",
split="train",
)
# Python の documents を読み込む
documents = load_dataset(
"Shuu12121/owl_code_search_hard_negative_datasets-Pre_kd",
name="documents_python",
split="train",
)ハードネガティブの抽出
# クエリ・文書テキストの辞書を構築
query_texts = dict(zip(queries["query_id"], queries["query"]))
doc_texts = dict(zip(documents["document_id"], documents["document"]))
# 閾値の設定
nv_threshold = 0.99 # positive スコアの 99% 未満をネガティブとする
# 1 サンプルの処理例
sample = scores[0]
query_text = query_texts[sample["query_id"]]
positive_doc = doc_texts[sample["document_ids"][0]] # scores[0] が正例
positive_score = sample["scores"][0]
hard_negatives = []
for doc_id, score in zip(sample["document_ids"][1:], sample["scores"][1:]):
if score < nv_threshold * positive_score and score != -1:
hard_negatives.append(doc_texts[doc_id])
print(f"Query: {query_text[:100]}...")
print(f"Positive: {positive_doc[:100]}...")
print(f"Hard negatives: {len(hard_negatives)}")