Spaces:
Running on Zero
Running on Zero
File size: 7,846 Bytes
0e50a2e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 | import html
import os
import re
import time
from concurrent.futures import ThreadPoolExecutor
from functools import lru_cache
from dotenv import load_dotenv
from openai import OpenAI
from pinecone import Pinecone
import fts_queries as fq
load_dotenv()
INDEX_NAME = "sec-fts"
NAMESPACE = "__default__"
EMBED_MODEL = "text-embedding-3-small"
TICKERS = ["aapl", "amzn", "f", "gm", "msft", "orcl"]
YEARS = list(range(2019, 2025))
INCLUDE_FIELDS = ["text", "ticker", "filing_type", "year", "chunk_index"]
RRF_K = 60
def missing_env() -> list[str]:
return [k for k in ("PINECONE_API_KEY", "OPENAI_API_KEY") if not os.environ.get(k)]
@lru_cache(maxsize=1)
def get_index():
return Pinecone(api_key=os.environ["PINECONE_API_KEY"]).index(name=INDEX_NAME)
@lru_cache(maxsize=1)
def get_openai() -> OpenAI:
return OpenAI(api_key=os.environ["OPENAI_API_KEY"])
@lru_cache(maxsize=512)
def _embed(text: str) -> tuple[float, ...]:
resp = get_openai().embeddings.create(model=EMBED_MODEL, input=[text])
return tuple(resp.data[0].embedding)
def embed(text: str) -> list[float]:
return list(_embed(text))
def metadata_filter(tickers: list[str], years: list[int]) -> dict | None:
filt = {}
if tickers:
filt["ticker"] = {"$in": list(tickers)}
if years:
filt["year"] = {"$in": [int(y) for y in years]}
return filt or None
def and_filters(*filters: dict | None) -> dict | None:
parts = [f for f in filters if f]
if not parts:
return None
return parts[0] if len(parts) == 1 else {"$and": parts}
def search(req: dict):
return get_index().documents.search(**req)
def dense_request(query: str, top_k: int, filt: dict | None = None) -> dict:
req = {
"namespace": NAMESPACE,
"top_k": top_k,
"score_by": [{"type": "dense_vector", "field": "embedding", "values": embed(query)}],
"include_fields": INCLUDE_FIELDS,
}
if filt:
req["filter"] = filt
return req
def run_search_text(query: str, tickers: list, years: list, top_k: int):
req = {
"namespace": NAMESPACE,
"top_k": top_k,
"score_by": [{"type": "text", "field": "text", "query": query}],
"include_fields": INCLUDE_FIELDS,
}
filt = metadata_filter(tickers, years)
if filt:
req["filter"] = filt
return search(req)
def run_search_semantic(query: str, tickers: list, years: list, top_k: int):
return search(dense_request(query, top_k, metadata_filter(tickers, years)))
def run_search_hybrid(semantic_query: str, text_filter: str, tickers: list, years: list, top_k: int):
filt = metadata_filter(tickers, years) or {}
if text_filter.strip():
filt["text"] = {"$match_all": text_filter.strip()}
return search(dense_request(semantic_query, top_k, filt or None))
def term_pattern(terms: list[str]) -> re.Pattern | None:
if not terms:
return None
stems = [re.escape(t[:-2] if len(t) > 5 else t) for t in terms]
return re.compile(r"\b(" + "|".join(stems) + r")\w*", re.IGNORECASE)
def highlight(text: str, terms: list[str]) -> tuple[str, str]:
escaped = html.escape(text)
pattern = term_pattern(terms)
if not pattern:
return escaped[:300] + ("..." if len(escaped) > 300 else ""), escaped
marked = pattern.sub(lambda m: f"<mark>{m.group(0)}</mark>", escaped)
first = pattern.search(escaped)
start = max(0, first.start() - 150) if first else 0
snippet_raw = escaped[start : start + 450]
snippet = pattern.sub(lambda m: f"<mark>{m.group(0)}</mark>", snippet_raw)
snippet = ("..." if start else "") + snippet + ("..." if start + 450 < len(escaped) else "")
return snippet, marked
def keyword_coverage(text: str, terms: list[str]) -> tuple[int, int]:
hits = sum(1 for t in terms if term_pattern([t]).search(text))
return hits, len(terms)
def match_record(m) -> dict:
d = m.to_dict()
return {
"id": getattr(m, "_id", d.get("_id")),
"score": getattr(m, "_score", getattr(m, "score", d.get("_score", 0.0))),
"ticker": str(d.get("ticker", "?")).upper(),
"year": int(d.get("year", 0)),
"filing": str(d.get("filing_type", "")).upper(),
"chunk": int(d["chunk_index"]) if d.get("chunk_index") is not None else "?",
"text": str(d.get("text", "")),
}
def _timed(fn):
start = time.perf_counter()
result = fn()
return result, (time.perf_counter() - start) * 1000
def run_compare(query: str, fts_state: dict, base_filter: dict | None, dense_extra_filter: dict | None, top_k: int) -> dict:
fts_req = fq.build_request(fts_state, top_k, INCLUDE_FIELDS)
fts_req["filter"] = and_filters(fts_req.get("filter"), base_filter)
if not fts_req["filter"]:
del fts_req["filter"]
def dense():
_, embed_ms = _timed(lambda: embed(query))
resp, search_ms = _timed(lambda: search(dense_request(query, top_k, and_filters(base_filter, dense_extra_filter))))
return resp, embed_ms, search_ms
with ThreadPoolExecutor(max_workers=2) as pool:
dense_future = pool.submit(dense)
fts_future = pool.submit(_timed, lambda: search(fts_req))
dense_resp, embed_ms, dense_ms = dense_future.result()
fts_resp, fts_ms = fts_future.result()
return {
"dense": [match_record(m) for m in dense_resp.matches],
"fts": [match_record(m) for m in fts_resp.matches],
"embed_ms": embed_ms,
"dense_ms": dense_ms,
"fts_ms": fts_ms,
"fts_req": fts_req,
"terms": fq.highlight_terms(fts_state),
}
def compare_stats(res: dict) -> dict:
dense, fts, terms = res["dense"], res["fts"], res["terms"]
dense_rank = {r["id"]: i for i, r in enumerate(dense, 1)}
fts_rank = {r["id"]: i for i, r in enumerate(fts, 1)}
shared = dense_rank.keys() & fts_rank.keys()
union = dense_rank.keys() | fts_rank.keys()
def avg_coverage(rows):
if not rows or not terms:
return None
return sum(keyword_coverage(r["text"], terms)[0] / len(terms) for r in rows) / len(rows)
return {
"dense_rank": dense_rank,
"fts_rank": fts_rank,
"shared": shared,
"union": union,
"by_id": {r["id"]: r for r in dense + fts},
"cov_dense": avg_coverage(dense),
"cov_fts": avg_coverage(fts),
"k": max(len(dense), len(fts)),
}
def rank_rows(res: dict, stats: dict) -> list[dict]:
terms = res["terms"]
rows = []
for doc_id in stats["union"]:
r = stats["by_id"][doc_id]
d, f = stats["dense_rank"].get(doc_id), stats["fts_rank"].get(doc_id)
hits, n = keyword_coverage(r["text"], terms)
rows.append({
"ticker": r["ticker"],
"year": r["year"],
"chunk": r["chunk"],
"dense rank": d,
"full-text rank": f,
"ฮ rank (dense โ FTS)": d - f if d and f else None,
"found by": "both" if d and f else ("dense" if d else "full-text"),
"keywords": f"{hits}/{n}" if n else "",
"snippet": r["text"][:140].replace("\n", " "),
})
rows.sort(key=lambda x: min(x["dense rank"] or 999, x["full-text rank"] or 999))
return rows
def rrf_fuse(stats: dict) -> list[tuple[str, float]]:
fused = {}
for ranks in (stats["dense_rank"], stats["fts_rank"]):
for doc_id, rank in ranks.items():
fused[doc_id] = fused.get(doc_id, 0.0) + 1 / (RRF_K + rank)
return sorted(fused.items(), key=lambda kv: kv[1], reverse=True)[: stats["k"]]
def rank_badge(doc_id: str, other_ranks: dict, other_name: str) -> str:
if doc_id in other_ranks:
return f"๐ also #{other_ranks[doc_id]} in {other_name}"
return f"โ only in {'dense' if other_name == 'full-text' else 'full-text'}"
|