You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.

407 lines
13 KiB
Python

"""Quick probe for free attention data on snapshot events.
This is a research helper, not a production pipeline.
Current sources:
- Wikimedia pageviews: scalable attention proxy
- GDELT Doc API: optional spot-check article count for a few names
"""
from __future__ import annotations
import argparse
import asyncio
import datetime as dt
import re
import time
from dataclasses import dataclass
from pathlib import Path
import pandas as pd
import requests
from sqlalchemy import text
from sqlalchemy.ext.asyncio import create_async_engine
from libs.common.config import get_settings
USER_AGENT = "codex-fithia2-free-attention-probe/1.0"
WIKI_SEARCH_URL = "https://en.wikipedia.org/w/api.php"
WIKI_PAGEVIEWS_URL = (
"https://wikimedia.org/api/rest_v1/metrics/pageviews/per-article/"
"en.wikipedia.org/all-access/all-agents/{article}/daily/{start}/{end}"
)
GDELT_DOC_URL = "https://api.gdeltproject.org/api/v2/doc/doc"
STOPWORDS = {
"inc",
"incorporated",
"corp",
"corporation",
"ltd",
"holdings",
"group",
"co",
"company",
"plc",
"nv",
}
MANUAL_WIKI_TITLE = {
"AMERICAN AIRLINES GROUP INC.": "American Airlines Group",
"ASTRONICS CORPORATION": "Astronics",
"CENTURY ALUMINUM COMPANY": "Century Aluminum",
"CLEANSPARK, INC.": "CleanSpark",
"CCC INTELLIGENT SOLUTIONS HOLDINGS INC.": "CCC Intelligent Solutions",
"FLUENCE ENERGY, INC.": "Fluence Energy",
"IMMUNITYBIO, INC.": "ImmunityBio",
"IMMUNITYBIO,\xa0INC.": "ImmunityBio",
"LYFT, INC.": "Lyft",
"MIRION TECHNOLOGIES, INC.": "Mirion Technologies",
"MOSAIC CO": "The Mosaic Company",
"NORWEGIAN CRUISE LINE HOLDINGS LTD.": "Norwegian Cruise Line Holdings",
"PAR PACIFIC HOLDINGS, INC.": "Par Pacific Holdings",
"PATTERSON-UTI ENERGY, INC.": "Patterson-UTI Energy",
"RITHM CAPITAL CORP.": "Rithm Capital",
"SOUNDHOUND AI, INC.": "SoundHound AI",
"TANGO THERAPEUTICS, INC.": "Tango Therapeutics",
"TERNS PHARMACEUTICALS, INC.": "Terns Pharmaceuticals",
"UNITY SOFTWARE INC.": "Unity Technologies",
}
@dataclass(frozen=True)
class ProbeRow:
ticker: str
issuer_name: str
event_date: dt.date
reaction_day_return: float
fwd_return_3d: float
fwd_return_5d: float
def _session() -> requests.Session:
s = requests.Session()
s.headers.update({"User-Agent": USER_AGENT})
return s
def _clean_tokens(text_value: str) -> list[str]:
tokens = re.findall(r"[A-Za-z0-9]+", text_value.lower().replace("\xa0", " "))
return [token for token in tokens if token not in STOPWORDS]
def _candidate_names(name: str) -> list[str]:
clean = name.replace("\xa0", " ").strip()
manual = MANUAL_WIKI_TITLE.get(clean.upper())
candidates = [candidate for candidate in [manual, clean] if candidate]
token_name = " ".join(_clean_tokens(clean))
if token_name:
candidates.append(token_name)
seen: set[str] = set()
result: list[str] = []
for candidate in candidates:
stripped = candidate.strip(" ,.")
if stripped and stripped not in seen:
result.append(stripped)
seen.add(stripped)
return result
def _title_match_score(name: str, title: str) -> float:
name_tokens = set(_clean_tokens(name))
title_tokens = set(_clean_tokens(title))
if not name_tokens or not title_tokens:
return 0.0
overlap = len(name_tokens & title_tokens)
score = overlap / max(1, len(name_tokens))
first_word = name.split()[0].lower() if name.split() else ""
if first_word and title.lower().startswith(first_word):
score += 0.1
return score
def resolve_wikipedia_title(session: requests.Session, issuer_name: str) -> tuple[str | None, float]:
best_score = 0.0
best_title: str | None = None
for candidate in _candidate_names(issuer_name):
response = session.get(
WIKI_SEARCH_URL,
params={
"action": "query",
"list": "search",
"srsearch": candidate,
"format": "json",
"srlimit": 5,
},
timeout=20,
)
response.raise_for_status()
hits = response.json().get("query", {}).get("search", [])
for hit in hits:
title = hit["title"]
score = _title_match_score(candidate, title)
if score > best_score:
best_score = score
best_title = title
if best_score >= 0.55:
return best_title, best_score
return None, best_score
def fetch_pageview_spike(
session: requests.Session,
article_title: str,
event_date: dt.date,
) -> dict[str, float] | None:
start = (event_date - dt.timedelta(days=20)).strftime("%Y%m%d")
end = (event_date + dt.timedelta(days=2)).strftime("%Y%m%d")
response = session.get(
WIKI_PAGEVIEWS_URL.format(
article=article_title.replace(" ", "_"),
start=start,
end=end,
),
timeout=20,
)
if response.status_code != 200:
return None
items = response.json().get("items", [])
if len(items) < 8:
return None
views = pd.DataFrame(
[(pd.to_datetime(item["timestamp"][:8]), item["views"]) for item in items],
columns=["date", "views"],
).sort_values("date")
event_ts = pd.Timestamp(event_date)
pre_event = views.loc[views["date"] < event_ts, "views"]
event_views = views.loc[views["date"] == event_ts, "views"]
if pre_event.empty or event_views.empty:
return None
baseline = float(pre_event.tail(10).median())
if baseline <= 0:
return None
event_value = float(event_views.iloc[0])
return {
"event_views": event_value,
"baseline_views": baseline,
"pageview_spike": event_value / baseline,
}
def fetch_gdelt_article_count(
session: requests.Session,
issuer_name: str,
event_date: dt.date,
) -> int | None:
exact_name = MANUAL_WIKI_TITLE.get(issuer_name.upper(), issuer_name)
phrase = exact_name.replace("\xa0", " ").replace('"', "")
if len(phrase) < 6:
return None
time.sleep(6.0)
start = (event_date - dt.timedelta(days=1)).strftime("%Y%m%d") + "000000"
end = (event_date + dt.timedelta(days=1)).strftime("%Y%m%d") + "235959"
response = session.get(
GDELT_DOC_URL,
params={
"query": f'"{phrase}"',
"mode": "ArtList",
"maxrecords": 50,
"format": "json",
"startdatetime": start,
"enddatetime": end,
},
timeout=30,
)
if response.status_code != 200:
return None
payload = response.json()
articles = payload.get("articles", [])
return len(articles)
async def load_probe_rows(snapshot_path: Path, split_name: str, limit: int) -> list[ProbeRow]:
snapshot = pd.read_parquet(snapshot_path)
snapshot = snapshot[
["event_id", "event_date", "reaction_day_return", "fwd_return_3d", "fwd_return_5d"]
].copy()
snapshot["event_date"] = pd.to_datetime(snapshot["event_date"]).dt.date
engine = create_async_engine(get_settings().postgres_dsn)
try:
async with engine.connect() as conn:
result = await conn.execute(
text(
"""
select e.event_id, e.event_type, sm.ticker, i.issuer_name
from events e
left join issuer_master i on i.issuer_id = e.issuer_id
left join symbol_master sm on sm.symbol_id = e.symbol_id
where e.event_id = any(:ids)
"""
),
{"ids": snapshot["event_id"].tolist()},
)
meta = pd.DataFrame(result.fetchall(), columns=result.keys())
finally:
await engine.dispose()
merged = snapshot.merge(meta, on="event_id", how="left")
merged = merged[merged["event_type"] == "earnings_release"].copy()
generic = (
(merged["issuer_name"].fillna("") == merged["ticker"].fillna("") + " Corporation")
| (merged["issuer_name"].fillna("") == merged["ticker"].fillna("") + " Inc.")
| (merged["issuer_name"].fillna("") == merged["ticker"].fillna("") + " Ltd.")
)
filtered = merged.loc[~generic & merged["issuer_name"].notna()].sort_values("event_date")
if limit > 0:
filtered = filtered.head(limit)
return [
ProbeRow(
ticker=str(row["ticker"]),
issuer_name=str(row["issuer_name"]),
event_date=row["event_date"],
reaction_day_return=float(row["reaction_day_return"]),
fwd_return_3d=float(row["fwd_return_3d"]),
fwd_return_5d=float(row["fwd_return_5d"]),
)
for _, row in filtered.iterrows()
]
def run_probe(
rows: list[ProbeRow],
output_csv: Path | None,
gdelt_limit: int,
) -> pd.DataFrame:
session = _session()
resolved_rows: list[dict[str, object]] = []
for idx, row in enumerate(rows):
title, match_score = resolve_wikipedia_title(session, row.issuer_name)
if not title:
continue
pageviews = fetch_pageview_spike(session, title, row.event_date)
if not pageviews:
continue
gdelt_count = None
if idx < gdelt_limit:
gdelt_count = fetch_gdelt_article_count(session, row.issuer_name, row.event_date)
signed_cont_3d = (1.0 if row.reaction_day_return >= 0 else -1.0) * row.fwd_return_3d
signed_cont_5d = (1.0 if row.reaction_day_return >= 0 else -1.0) * row.fwd_return_5d
resolved_rows.append(
{
"ticker": row.ticker,
"issuer_name": row.issuer_name,
"article_title": title,
"match_score": round(match_score, 3),
"event_date": row.event_date.isoformat(),
"reaction_day_return": row.reaction_day_return,
"fwd_return_3d": row.fwd_return_3d,
"fwd_return_5d": row.fwd_return_5d,
"signed_cont_3d": signed_cont_3d,
"signed_cont_5d": signed_cont_5d,
**pageviews,
"gdelt_article_count_3d": gdelt_count,
}
)
df = pd.DataFrame(resolved_rows)
if output_csv and not df.empty:
output_csv.parent.mkdir(parents=True, exist_ok=True)
df.to_csv(output_csv, index=False)
return df
def print_summary(df: pd.DataFrame, sampled_rows: int) -> None:
print(f"resolved_rows={len(df)} sampled_rows={sampled_rows}")
if df.empty:
return
preview_cols = [
"ticker",
"issuer_name",
"article_title",
"match_score",
"pageview_spike",
"signed_cont_3d",
"signed_cont_5d",
"gdelt_article_count_3d",
]
print(df[preview_cols].to_string(index=False))
median_spike = float(df["pageview_spike"].median())
high = df[df["pageview_spike"] >= median_spike]
low = df[df["pageview_spike"] < median_spike]
print("")
print(f"median_pageview_spike={median_spike:.3f}")
print(f"high_group_n={len(high)} low_group_n={len(low)}")
print(f"high_signed_cont_3d_mean={high['signed_cont_3d'].mean():.4f}")
print(f"low_signed_cont_3d_mean={low['signed_cont_3d'].mean():.4f}")
print(f"high_signed_cont_5d_mean={high['signed_cont_5d'].mean():.4f}")
print(f"low_signed_cont_5d_mean={low['signed_cont_5d'].mean():.4f}")
print(f"corr(pageview_spike,signed_cont_3d)={df['pageview_spike'].corr(df['signed_cont_3d']):.4f}")
print(f"corr(pageview_spike,signed_cont_5d)={df['pageview_spike'].corr(df['signed_cont_5d']):.4f}")
gdelt = df["gdelt_article_count_3d"].dropna()
if not gdelt.empty:
print(
"gdelt_counts_sample="
+ ", ".join(str(int(value)) for value in gdelt.tolist())
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Probe free attention data against snapshot outcomes.")
parser.add_argument(
"--snapshot",
default="data/datasets/snapshots/midcap-filtered/test.parquet",
help="Snapshot parquet path",
)
parser.add_argument(
"--split",
default="test",
help="Label only for reporting; snapshot path determines actual data",
)
parser.add_argument(
"--limit",
type=int,
default=30,
help="Number of filtered earnings events to sample",
)
parser.add_argument(
"--gdelt-limit",
type=int,
default=5,
help="How many resolved rows to spot-check with GDELT article counts",
)
parser.add_argument(
"--output-csv",
default="data/research/free_attention_probe_test_sample.csv",
help="Output CSV path",
)
return parser.parse_args()
async def main() -> None:
args = parse_args()
snapshot_path = Path(args.snapshot)
rows = await load_probe_rows(snapshot_path, args.split, args.limit)
df = run_probe(rows, Path(args.output_csv), args.gdelt_limit)
print_summary(df, sampled_rows=len(rows))
if not df.empty:
print(f"saved_csv={args.output_csv}")
if __name__ == "__main__":
asyncio.run(main())