/
luvtrippin
/
RAG-Diploma
Обзор
Документация
Войти
/
luvtrippin
/
RAG-Diploma
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
src/loader.py
308 строк
9 KB
Timur
Prepare reproducible RAG experiment project
25 июн 2026, 18:20
25 июн 2026, 18:20
b279d80
Код
Авторство
О чём код?
import json from langchain_core.documents import Document SCIFACT_LABEL_TO_DECISION = { "SUPPORT": "yes", "CONTRADICT": "no", "NOT_ENOUGH_INFO": "maybe", } def load_custom_dataset(path, limit=None): with open(path, "r", encoding="utf-8") as f: data = json.load(f) documents = [] qa_pairs = [] for i, (qid, item) in enumerate(data.items()): if qid.startswith("_"): continue if limit is not None and len(qa_pairs) >= limit: break question = item["QUESTION"] contexts = item["CONTEXTS"] if isinstance(contexts, str): contexts = [contexts] answer = item["ANSWER"] qa_pairs.append({ "id": qid, "question": question, "contexts": contexts, "answer": answer, "long_answer": answer, "relevant_contexts": len(contexts), "answer_format": "full", }) for context_index, ctx in enumerate(contexts): documents.append( Document( page_content=ctx, metadata={ "question_id": qid, "context_index": context_index, "source": "custom", }, ) ) return documents, qa_pairs def load_sciq_dataset(path, limit=None): with open(path, "r", encoding="utf-8") as f: data = json.load(f) if isinstance(data, dict): rows = data.get("items") or [ {**item, "id": qid} for qid, item in data.items() if not str(qid).startswith("_") ] else: rows = data documents = [] qa_pairs = [] for index, item in enumerate(rows): if limit is not None and len(qa_pairs) >= limit: break support = str(item.get("support") or "").strip() if not support: continue qid = str(item.get("id") or item.get("question_id") or f"sciq_{index:05d}") question = str(item["question"]).strip() answer = str(item["correct_answer"]).strip() qa_pairs.append({ "id": qid, "question": question, "contexts": [support], "answer": answer, "long_answer": answer, "relevant_contexts": 1, "answer_format": "open_short", "choices": item.get("choices") or [], }) documents.append( Document( page_content=support, metadata={ "question_id": qid, "context_index": 0, "source": "sciq", }, ) ) return documents, qa_pairs def _researchqa_reference_texts(item): texts = [] for section_index, reference in enumerate(item.get("expected_references") or []): alternatives = reference.get("alternatives") or [] if not alternatives: continue text = str(alternatives[0]).strip() if text: texts.append({ "section_index": section_index, "section_label": reference.get("section_label", ""), "text": text, }) return texts def load_researchqa_dataset(path, limit=None): with open(path, "r", encoding="utf-8") as f: data = json.load(f) rows = data.get("items", data if isinstance(data, list) else []) documents = [] qa_pairs = [] for index, item in enumerate(rows): if limit is not None and len(qa_pairs) >= limit: break qid = str(item.get("id") or item.get("row_id") or f"researchqa_{index:05d}") question = str(item.get("question") or "").strip() answer = str(item.get("expected_answer") or "").strip() references = _researchqa_reference_texts(item) if not question or not answer or not references: continue relevant_doc_ids = set() contexts = [] for reference in references: doc_id = f"{qid}::section_{reference['section_index']}" relevant_doc_ids.add(doc_id) contexts.append(reference["text"]) documents.append( Document( page_content=reference["text"], metadata={ "doc_id": doc_id, "question_id": qid, "context_index": reference["section_index"], "section_label": reference["section_label"], "source": "researchqa", "domain": item.get("domain", ""), "question_type": item.get("question_type", ""), "paper_id": item.get("paper_id", ""), }, ) ) qa_pairs.append({ "id": qid, "question": question, "contexts": contexts, "answer": answer, "long_answer": answer, "relevant_contexts": len(relevant_doc_ids), "relevant_doc_ids": relevant_doc_ids, "answer_format": "full", "question_type": item.get("question_type", ""), "domain": item.get("domain", ""), "judge_rubric": item.get("judge_rubric", ""), "expected_refusal": item.get("expected_refusal"), }) return documents, qa_pairs def load_pubmedqa(path, limit=None): with open(path, "r", encoding="utf-8") as f: data = json.load(f) documents = [] qa_pairs = [] for i, (qid, item) in enumerate(data.items()): if limit is not None and i >= limit: break question = item["QUESTION"] contexts = item["CONTEXTS"] answer = item["final_decision"] long_answer = item.get("LONG_ANSWER", "") qa_pairs.append({ "id": qid, "question": question, "contexts": contexts, "answer": answer, "long_answer": long_answer, "relevant_contexts": len(contexts), }) for context_index, ctx in enumerate(contexts): documents.append( Document( page_content=ctx, metadata={ "question_id": qid, "context_index": context_index, "source": "pubmedqa" } ) ) return documents, qa_pairs def _read_jsonl(path): rows = [] with open(path, "r", encoding="utf-8") as f: for line in f: line = line.strip() if line: rows.append(json.loads(line)) return rows def _scifact_doc_text(item): title = item.get("title") or "" abstract = item.get("abstract") or [] if isinstance(abstract, list): abstract_text = " ".join(abstract) else: abstract_text = str(abstract) return f"{title}\n{abstract_text}".strip() def _scifact_claim_label(claim): labels = [] for evidence_items in (claim.get("evidence") or {}).values(): for item in evidence_items: label = item.get("label") if label: labels.append(label) if "SUPPORT" in labels: return "SUPPORT" if "CONTRADICT" in labels: return "CONTRADICT" return "NOT_ENOUGH_INFO" def load_scifact(corpus_path, claims_path, limit=None, index_cited_only=False): corpus = _read_jsonl(corpus_path) corpus_by_id = {int(item["doc_id"]): item for item in corpus} claims = _read_jsonl(claims_path) cited_pool = set() if index_cited_only: for claim in claims: cited_pool.update(int(doc_id) for doc_id in claim.get("cited_doc_ids", [])) documents = [] for item in corpus: doc_id = int(item["doc_id"]) if index_cited_only and doc_id not in cited_pool: continue documents.append( Document( page_content=_scifact_doc_text(item), metadata={ "doc_id": doc_id, "source": "scifact", "title": item.get("title", ""), }, ) ) qa_pairs = [] for claim in claims: if limit is not None and len(qa_pairs) >= limit: break cited_doc_ids = [int(doc_id) for doc_id in claim.get("cited_doc_ids", [])] label = _scifact_claim_label(claim) contexts = [ _scifact_doc_text(corpus_by_id[doc_id]) for doc_id in cited_doc_ids if doc_id in corpus_by_id ] qa_pairs.append({ "id": str(claim["id"]), "question": claim["claim"], "contexts": contexts, "answer": SCIFACT_LABEL_TO_DECISION[label], "original_label": label, "long_answer": label, "relevant_contexts": len(set(cited_doc_ids)), "relevant_doc_ids": set(cited_doc_ids), }) return documents, qa_pairs