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
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())
|