diff --git a/.gitignore b/.gitignore deleted file mode 100644 index 9866db2..0000000 --- a/.gitignore +++ /dev/null @@ -1,4 +0,0 @@ -embedding_firehose_client -*.jsonl -.venv/ -__pycache__/ diff --git a/.python-version b/.python-version deleted file mode 100644 index 6324d40..0000000 --- a/.python-version +++ /dev/null @@ -1 +0,0 @@ -3.14 diff --git a/README.md b/README.md index f8dda0b..5c5f8a6 100644 --- a/README.md +++ b/README.md @@ -1,106 +1,16 @@ -# Divepool Embedding Firehose Client (experimental) +# Divepool Embedding Firehose Client — retired -Example client for the Divepool embedding firehose and search API — a real-time stream of 128d [EmbeddingGemma](https://huggingface.co/google/embeddinggemma-300m) embeddings (Matryoshka truncation, L2-normalized) for Bluesky posts and profiles, plus a semantic search endpoint. Both the API and this client are **highly experimental** and subject to change. +**API access is deactivated until further notice.** -## Stream client +The public REST data endpoints and the embedding firehose this client targeted have +been retired. The example code that lived here has been removed. -Real-time embedding stream. Public clients receive 128d vectors; token-authenticated clients receive full 768d. +Agent access to divepool's discovery continues via **MCP**: -```bash -go run . stream # public endpoint (128d) -go run . stream -token x # with bearer token (768d) -go run . stream -token x -out events.jsonl # write events to file -``` +- MCP endpoint: https://divepool.com/mcp +- Agent front door: https://divepool.com/llms.txt -Without a subcommand, defaults to `stream` for backwards compatibility. - -## Search client - -Semantic search across all indexed Bluesky posts. Bearer token optional — non-bearer gets 128d embeddings, bearer gets 768d. - -```bash -# Global semantic search -go run . search "machine learning" # public search (top 20) -go run . search -token x "machine learning" # bearer search (768d embeddings) -go run . search -limit 100 "machine learning" # more results (max 400 global) -go run . search -distinct=false "machine learning" # multiple posts per account -go run . search -cluster "machine learning" # with topic clusters -go run . search -embeddings "machine learning" # include per-result embeddings -go run . search -langs de,en "machine learning" # only these languages -go run . search -since 2026-06-01T00:00:00Z "machine learning" # only recent posts (-until for an upper bound) -go run . search -min-score 0.55 "machine learning" # drop weak matches (cosine = -score) - -# Account-scoped (max 1200) -go run . search -did did:plc:abc123 # browse account (newest, clustered) -go run . search -did did:plc:abc123 "machine learning" # search within account -go run . search -did did:plc:abc123 -cluster "machine learning" # search within account + cluster - -# Search by example post (within that account — global search-by-example -# requires round-tripping the post text through a query, see experiments/amplifier_gap.py) -go run . search -did did:plc:abc123 -rkey 3abc # find similar posts in the account -go run . search -did did:plc:abc123 -rkey 3abc -cluster # find similar + cluster - -# Cluster a set of accounts by their topical medoids (max 500 DIDs) -# Each account contributes its top 3 medoids → equal voice per account. -# Returns medoid posts in `results` (tagged with cluster_id/topics) and -# HDBSCAN groupings in `clusters`. Use it to find topical groups inside -# any DID set you can produce — e.g. your Bluesky followers. -go run . search -dids did:plc:a,did:plc:b,did:plc:c # ad-hoc -go run . search -dids-file followers.txt # one DID per line -# Optional: -viewer-did scopes the clustering to the languages that account -# posts in (derived from its medoids server-side; echoed back as `langs`) — -# without it, multilingual follower sets produce clusters in every language. -go run . search -dids-file followers.txt -viewer-did did:plc:abc123 # language-scoped -``` - -## Medoids client - -Fetch top cluster medoids (representative posts) for accounts. Bearer token optional — non-bearer gets 128d embeddings, bearer gets 768d. - -```bash -go run . medoids did:plc:abc123 did:plc:def456 # top 3 medoids per account (128d) -go run . medoids -token x did:plc:abc123 # bearer (768d) -``` - -## Similar accounts client - -Rank the semantically nearest accounts to a given account, corpus-wide — the -account-level search-by-example that `search -rkey` (account-scoped) doesn't -cover. Each result carries `score` (cross-topic affinity), `margin` (score -minus the candidate-pool mean — the discriminative number; ≤ 0 is noise), and -the matched medoid post explaining *why* the account is similar. No token -needed (the response carries no embeddings). - -```bash -go run . similar did:plc:abc123 # nearest accounts (auto-scoped to the account's languages) -go run . similar -limit 50 did:plc:abc123 # more results (server max 100) -go run . similar -min-cluster-size 5 did:plc:abc123 # suppress one-off topic clusters on both sides -go run . similar -langs de did:plc:abc123 # only German topic matches -go run . similar -langs all did:plc:abc123 # cross-lingual neighbors -``` - -### API documentation - -Full API spec (search endpoints, medoids, firehose protocol, schemas): **[OpenAPI 3.1](https://divepool.social/api/v1/openapi.json)** - -## Experiments - -Python scripts in `experiments/` exploring personalized feeds, trend detection, and topic analysis on top of the firehose and search API. Even more heavily vibe-coded than the rest of this repo. Each script's docstring describes usage. - -```bash -pip install -r experiments/requirements.txt -``` - -## Verify embeddings - -Verify events from the stream against local EmbeddingGemma output: - -```bash -python3 -m venv .venv && .venv/bin/pip install -r requirements.txt -HF_TOKEN=hf_... .venv/bin/python3 verify.py events.jsonl -``` - -Takes JSONL captured with `-out` (`did`, `col`, `rkey`, `lang`, `c`, `r`). Fetches each record via Bluesky API, prepares text identically to Divepool (posts: text + image/video alt texts + tags; profiles: displayName + description), embeds locally, and compares via cosine similarity. Requires a [Hugging Face token](https://huggingface.co/settings/tokens) for the gated model. +This may change again in the future. ## License diff --git a/batch.go b/batch.go deleted file mode 100644 index 729ad11..0000000 --- a/batch.go +++ /dev/null @@ -1,29 +0,0 @@ -package main - -import "fmt" - -// Batch is the columnar wire format for the embedding firehose. -// Each field is a parallel array — index i across all fields describes one text. -// -// Embeddings are EmbeddingGemma-300m (Matryoshka) float32 vectors. Two per text: -// - C (cluster): for clustering/similarity analysis -// - R (retrieval): for semantic search -// -// Open clients receive 128d (L2-normalized Matryoshka truncation). -// Token-authenticated clients receive full 768d. -type Batch struct { - DID []string `json:"did"` // AT Protocol DID (e.g. "did:plc:abc123") - Col []string `json:"col"` // Collection NSID (e.g. "app.bsky.feed.post") - Rkey []string `json:"rkey"` // Record key - Lang []string `json:"lang"` // Detected language ("en", "de") - C [][]float32 `json:"c"` // Cluster embeddings (128d open, 768d token) - R [][]float32 `json:"r"` // Retrieval embeddings (128d open, 768d token) -} - -// Len returns the number of items in the batch. -func (b *Batch) Len() int { return len(b.DID) } - -// ATURI returns the AT URI for item i. -func (b *Batch) ATURI(i int) string { - return fmt.Sprintf("at://%s/%s/%s", b.DID[i], b.Col[i], b.Rkey[i]) -} diff --git a/experiments/amplifier_gap.py b/experiments/amplifier_gap.py deleted file mode 100644 index 4bd00ad..0000000 --- a/experiments/amplifier_gap.py +++ /dev/null @@ -1,260 +0,0 @@ -#!/usr/bin/env python3 -""" -Amplifier Gap — who should be reposting you but isn't. - -Reposts are the reach multiplier on Bluesky, so the warmest growth targets are -high-reach accounts that are semantically close to what you post but don't yet -engage with you. This tool finds them: - -1. Divepool /similar-accounts: the accounts nearest to your topic medoids - across the whole indexed network, one call. Margin-above-baseline scoring - is baked in server-side — this script used to round-trip each medoid post's - text through a global /search and mean-center per beat client-side; the - endpoint now does both (margin = score above the candidate pool's mean). -2. Your current amplifiers from the Bluesky API (getRepostedBy + getLikes on - your recent posts) -> subtracted, along with accounts you already follow - unless --include-follows is set -3. Survivors enriched with profiles and ranked by alignment x reach - (margin x log10 followers). Each shows its matched post — the candidate's - own post nearest your topics — as the "why". - -Token optional and cosmetic here: /similar-accounts returns no embeddings, so -anonymous output is identical. - -Usage: - python3 amplifier_gap.py [--token TOKEN] [--top 30] -""" - -import argparse -import json -import math -import sys -import urllib.error -import urllib.parse -import urllib.request -from concurrent.futures import ThreadPoolExecutor - -DIVEPOOL = "https://divepool.social/api/v1" -BSKY_API = "https://public.api.bsky.app/xrpc" - -RECENT_POSTS = 50 # how many of your recent posts to scan for amplifiers -ENGAGED_POSTS = 20 # top posts (by engagement) whose likers/reposters count as amplifiers -SIMILAR_LIMIT = 100 # server max for /similar-accounts - - -# ── Divepool ───────────────────────────────────────────────────────────────── - -def divepool_post(path: str, payload: dict, token: str | None) -> dict: - headers = {"Content-Type": "application/json"} - if token: - headers["Authorization"] = f"Bearer {token}" - req = urllib.request.Request(f"{DIVEPOOL}{path}", data=json.dumps(payload).encode(), - headers=headers, method="POST") - with urllib.request.urlopen(req, timeout=30) as resp: - return json.loads(resp.read()) - - -# ── Bluesky public API ─────────────────────────────────────────────────────── - -def bsky_get(method: str, params: dict) -> dict | None: - qs = urllib.parse.urlencode(params) - req = urllib.request.Request(f"{BSKY_API}/{method}?{qs}", - headers={"Accept": "application/json"}) - try: - with urllib.request.urlopen(req, timeout=10) as resp: - return json.loads(resp.read()) - except urllib.error.HTTPError: - return None - - -def resolve_did(handle_or_did: str) -> str: - if handle_or_did.startswith("did:"): - return handle_or_did - data = bsky_get("com.atproto.identity.resolveHandle", {"handle": handle_or_did}) - if not data: - sys.exit(f"could not resolve handle {handle_or_did}") - return data["did"] - - -def get_profiles(dids: list[str]) -> dict[str, dict]: - """Batch-fetch profiles (25 per call).""" - profiles = {} - for i in range(0, len(dids), 25): - qs = "&".join(f"actors={urllib.parse.quote(d)}" for d in dids[i:i+25]) - req = urllib.request.Request(f"{BSKY_API}/app.bsky.actor.getProfiles?{qs}", - headers={"Accept": "application/json"}) - try: - with urllib.request.urlopen(req, timeout=10) as resp: - for p in json.loads(resp.read()).get("profiles", []): - profiles[p["did"]] = p - except urllib.error.HTTPError: - pass - return profiles - - -def get_post_texts(uris: list[str]) -> dict[str, str]: - """Batch-fetch post texts (25 per getPosts call). Deleted posts are absent.""" - texts = {} - for i in range(0, len(uris), 25): - qs = "&".join(f"uris={urllib.parse.quote(u)}" for u in uris[i:i+25]) - req = urllib.request.Request(f"{BSKY_API}/app.bsky.feed.getPosts?{qs}", - headers={"Accept": "application/json"}) - try: - with urllib.request.urlopen(req, timeout=10) as resp: - for p in json.loads(resp.read()).get("posts", []): - texts[p["uri"]] = p.get("record", {}).get("text", "") - except urllib.error.HTTPError: - pass - return texts - - -def recent_posts(did: str) -> list[dict]: - """Your recent original posts (no replies/reposts), with engagement counts.""" - posts, cursor = [], "" - while len(posts) < RECENT_POSTS: - params = {"actor": did, "limit": 100, "filter": "posts_no_replies"} - if cursor: - params["cursor"] = cursor - data = bsky_get("app.bsky.feed.getAuthorFeed", params) - if not data: - break - for item in data.get("feed", []): - post = item["post"] - if post["author"]["did"] != did: # skip reposts of others - continue - posts.append(post) - cursor = data.get("cursor") - if not cursor: - break - return posts[:RECENT_POSTS] - - -def engagers(post_uri: str) -> set[str]: - """DIDs that liked or reposted a post (first page of each — amplifier - detection, not a census).""" - dids = set() - likes = bsky_get("app.bsky.feed.getLikes", {"uri": post_uri, "limit": 100}) - for l in (likes or {}).get("likes", []): - dids.add(l["actor"]["did"]) - reposts = bsky_get("app.bsky.feed.getRepostedBy", {"uri": post_uri, "limit": 100}) - for r in (reposts or {}).get("repostedBy", []): - dids.add(r["did"]) - return dids - - -def follows(did: str) -> set[str]: - """Everyone `did` follows (paged).""" - out, cursor = set(), "" - while True: - params = {"actor": did, "limit": 100} - if cursor: - params["cursor"] = cursor - data = bsky_get("app.bsky.graph.getFollows", params) - if not data: - break - for f in data.get("follows", []): - out.add(f["did"]) - cursor = data.get("cursor") - if not cursor: - break - return out - - -# ── Main ───────────────────────────────────────────────────────────────────── - -def main(): - ap = argparse.ArgumentParser(description="Find accounts that should be reposting you but aren't.") - ap.add_argument("account", help="your handle or DID") - ap.add_argument("--token", default=None, help="Divepool bearer token (optional)") - ap.add_argument("--top", type=int, default=30, help="how many targets to show") - ap.add_argument("--min-cluster-size", type=int, default=0, - help="ignore topic clusters smaller than this, yours and theirs (noise filter)") - ap.add_argument("--include-follows", action="store_true", - help="keep accounts you already follow in the list") - args = ap.parse_args() - - did = resolve_did(args.account) - print(f"Amplifier gap for {args.account} ({did})") - print("─" * 72) - - # ── Your beats: medoids with topic labels (display only) ───────────── - medoids_resp = divepool_post("/medoids", {"dids": [did]}, args.token) - medoids = medoids_resp.get("accounts", {}).get(did, {}).get("medoids", []) - if medoids: - print(f"\nYour beats (top {len(medoids)} medoids):") - for m in medoids: - label = " · ".join(m.get("topics", [])) or "(unlabeled)" - print(f" [{m['cluster_id']}] {label} ({m['cluster_size']} posts)") - - # ── Semantic neighborhood: one /similar-accounts call ──────────────── - payload = {"did": did, "limit": SIMILAR_LIMIT} - if args.min_cluster_size: - payload["min_cluster_size"] = args.min_cluster_size - res = divepool_post("/similar-accounts", payload, args.token) - accounts = res.get("accounts", []) - if not accounts: - sys.exit("no similar accounts — this account isn't indexed in Divepool yet " - "(needs posts ≥120 chars with a detectable language), or " - "--min-cluster-size filtered all its topic clusters") - if res.get("langs"): - print(f"\nLanguage scope (derived from your posts): {', '.join(res['langs'])}") - # Server margins are centered on the full candidate pool's mean; ≤ 0 means - # at/below that account's own similarity baseline — noise, drop upfront. - candidates = {a["did"]: a for a in accounts if a["margin"] > 0} - print(f"Semantic neighborhood: {len(accounts)} accounts, {len(candidates)} decisively " - f"above baseline (margin > 0, pool mean {res.get('score_mean', 0.0):.2f})") - - # ── Who already engages ─────────────────────────────────────────────── - posts = recent_posts(did) - posts.sort(key=lambda p: p.get("repostCount", 0) + p.get("likeCount", 0), reverse=True) - scan = posts[:ENGAGED_POSTS] - already = {did} - with ThreadPoolExecutor(max_workers=4) as ex: - for s in ex.map(engagers, (p["uri"] for p in scan)): - already |= s - print(f"Current amplifiers (likers/reposters of your top {len(scan)} recent posts): {len(already) - 1}") - - followed = set() - if not args.include_follows: - followed = follows(did) - print(f"Accounts you follow (subtracted, --include-follows keeps them): {len(followed)}") - - gap = {d: a for d, a in candidates.items() if d not in already and d not in followed} - print(f"Gap: {len(gap)} semantically-near accounts that don't engage with you yet") - - # ── The why: each candidate's matched post (nearest to your topics) ─── - matched_uri = {d: f"at://{d}/{a['matched']['collection']}/{a['matched']['rkey']}" - for d, a in gap.items()} - matched_texts = get_post_texts(list(matched_uri.values())) - - # ── Enrich + rank: alignment × reach ───────────────────────────────── - profiles = get_profiles(list(gap)) - ranked = [] - for d, a in gap.items(): - p = profiles.get(d) - if not p: - continue - fl = p.get("followersCount", 0) - ranked.append((a["margin"] * math.log10(fl + 10), fl, a, p)) - ranked.sort(reverse=True, key=lambda r: r[0]) - - print(f"\nTop {min(args.top, len(ranked))} warm targets (margin above pool baseline × reach):") - print("─" * 72) - for rank_score, fl, a, p in ranked[:args.top]: - handle = p.get("handle", "?") - name = (p.get("displayName") or "").strip() - bio = " ".join((p.get("description") or "").split())[:90] - print(f"@{handle} ({fl:,} followers, margin +{a['margin']:.2f})") - if name: - print(f" {name}") - if bio: - print(f" {bio}") - text = matched_texts.get(matched_uri[a["did"]]) - if text: - snippet = " ".join(text.split())[:110] - print(f" their closest post ({a['matched']['cluster_size']}-post topic): “{snippet}”") - print() - - -if __name__ == "__main__": - main() diff --git a/experiments/for_you.py b/experiments/for_you.py deleted file mode 100644 index 94670a1..0000000 --- a/experiments/for_you.py +++ /dev/null @@ -1,505 +0,0 @@ -#!/usr/bin/env python3 -""" -For You — personalized Bluesky feed from the Divepool firehose. - -Uses ~15 medoid centroid embeddings from search API clusters. -Window-based selection: scores every firehose post, picks the best per window. -Includes detailed filter stats (for_you_stats.json). - -Usage: - python3 for_you.py --window 60 [--top 3] [--floor 0.5] -""" - -import argparse -import json -import os -import sys -import threading -import time -import urllib.error -import urllib.parse -import urllib.request -from collections import defaultdict - -import numpy as np - -DIVEPOOL_SEARCH = "https://divepool.social/api/v1/search" -DIVEPOOL_STREAM = "https://divepool.social/api/v1/embeddings" -BSKY_API = "https://public.api.bsky.app/xrpc" -STATS_FILE = os.path.join(os.path.dirname(os.path.abspath(__file__)), "for_you_stats.json") - - -# ── API helpers ────────────────────────────────────────────────────────────── - -def divepool_search(token, query, limit=100, did=None, cluster=False, distinct=True): - payload = {"query": query, "limit": limit, "distinct": distinct} - if did: - payload["did"] = did - if cluster: - payload["cluster"] = True - data = json.dumps(payload).encode() - req = urllib.request.Request(DIVEPOOL_SEARCH, data=data, headers={ - "Content-Type": "application/json", - "Authorization": f"Bearer {token}", - }, method="POST") - with urllib.request.urlopen(req, timeout=15) as resp: - return json.loads(resp.read()) - - -def resolve_post_text(did, collection, rkey): - uri = f"at://{did}/{collection}/{rkey}" - url = f"{BSKY_API}/app.bsky.feed.getPostThread?uri={urllib.parse.quote(uri)}&depth=0" - req = urllib.request.Request(url, headers={"Accept": "application/json"}) - try: - with urllib.request.urlopen(req, timeout=5) as resp: - data = json.loads(resp.read()) - post = data.get("thread", {}).get("post", {}) - text = post.get("record", {}).get("text", "") - handle = post.get("author", {}).get("handle", "") - return handle, text - except Exception: - return None, None - - -def bsky_link(handle, rkey): - return f"https://bsky.app/profile/{handle}/post/{rkey}" - - -def tid_to_timestamp(tid): - """Decode an AT Protocol TID (base32-sortable) to a Unix timestamp in seconds. - TIDs encode microseconds-since-epoch in the high 53 bits.""" - charset = "234567abcdefghijklmnopqrstuvwxyz" - try: - n = 0 - for ch in tid: - n = n * 32 + charset.index(ch) - # high 53 bits = microseconds since epoch - usec = n >> 10 - return usec / 1_000_000 - except (ValueError, IndexError): - return 0.0 - - -# ── Firehose streaming ────────────────────────────────────────────────────── - -def stream_firehose(token, post_queue, stop_event): - import zstandard as zstd - while not stop_event.is_set(): - try: - req = urllib.request.Request(DIVEPOOL_STREAM, headers={ - "Authorization": f"Bearer {token}", - }) - with urllib.request.urlopen(req, timeout=30) as resp: - dctx = zstd.ZstdDecompressor() - reader = dctx.stream_reader(resp) - buf = b"" - while not stop_event.is_set(): - chunk = reader.read(8192) - if not chunk: - break - buf += chunk - while b"\n" in buf: - line, buf = buf.split(b"\n", 1) - if not line.strip(): - continue - try: - batch = json.loads(line) - except json.JSONDecodeError: - continue - if not isinstance(batch, dict): - continue - dids = batch.get("did") - if not isinstance(dids, list) or not dids: - continue - cols = batch.get("col", []) - rkeys = batch.get("rkey", []) - langs = batch.get("lang", []) - cs = batch.get("c", []) - for i in range(len(dids)): - col = cols[i] if i < len(cols) else "" - if "feed.post" not in col: - continue - c_emb = cs[i] if i < len(cs) else [] - if not c_emb: - continue - rkey = rkeys[i] if i < len(rkeys) else "" - lang = langs[i] if i < len(langs) else "" - post_queue.append((dids[i], col, rkey, lang, c_emb)) - except Exception: - if not stop_event.is_set(): - time.sleep(2) - - -# ── Stats tracker ────────────────────────────────────────────────────────── - -class FilterStats: - """Tracks scoring distribution and selection stats.""" - - BINS = [0.0, 0.3, 0.4, 0.5, 0.55, 0.6, 0.65, 0.7, 0.75, 0.8, 0.85, 0.9, 0.95, 1.0] - - def __init__(self): - self.started = time.time() - self.total = 0 - self.skip_seen_self = 0 - self.fail_resolve = 0 - self.fail_spam = 0 - self.shown = 0 - self.windows_elapsed = 0 - self.windows_empty = 0 # windows where nothing was good enough - # histogram of composite scores for all scored posts - self.score_hist = np.zeros(len(self.BINS) - 1, dtype=np.int64) - # histogram of shown post scores - self.shown_hist = np.zeros(len(self.BINS) - 1, dtype=np.int64) - # per-medoid hit counts (best medoid for shown posts) - self.medoid_hits = defaultdict(int) - # track recent shown timestamps for rate calc - self._shown_times = [] - # recency: track age of incoming posts - self._recent_ages = [] - - def record_age(self, rkey): - ts = tid_to_timestamp(rkey) - if ts > 0: - age = time.time() - ts - self._recent_ages.append((time.time(), age)) - - def age_stats(self): - now = time.time() - cutoff = now - 300 - self._recent_ages = [(t, a) for t, a in self._recent_ages if t > cutoff] - if not self._recent_ages: - return None - ages = [a for _, a in self._recent_ages] - return { - "count": len(ages), - "median_sec": round(float(np.median(ages)), 1), - "p90_sec": round(float(np.percentile(ages, 90)), 1), - "max_sec": round(max(ages), 1), - "min_sec": round(min(ages), 1), - "pct_under_60s": round(100 * sum(1 for a in ages if a < 60) / len(ages), 1), - "pct_under_300s": round(100 * sum(1 for a in ages if a < 300) / len(ages), 1), - } - - def record_score(self, score): - idx = np.searchsorted(self.BINS, score, side="right") - 1 - idx = max(0, min(idx, len(self.score_hist) - 1)) - self.score_hist[idx] += 1 - - def record_shown(self, score): - self.shown += 1 - self._shown_times.append(time.time()) - idx = np.searchsorted(self.BINS, score, side="right") - 1 - idx = max(0, min(idx, len(self.shown_hist) - 1)) - self.shown_hist[idx] += 1 - - def shown_per_min(self): - now = time.time() - cutoff = now - 300 - self._shown_times = [t for t in self._shown_times if t > cutoff] - if not self._shown_times: - return 0.0 - window = now - self._shown_times[0] - if window < 10: - return 0.0 - return len(self._shown_times) / (window / 60) - - def _hist_to_dict(self, hist): - out = {} - for i, count in enumerate(hist): - if count > 0: - lo, hi = self.BINS[i], self.BINS[i + 1] - out[f"{lo:.2f}-{hi:.2f}"] = int(count) - return out - - def to_dict(self, floor): - elapsed = time.time() - self.started - scored = self.total - self.skip_seen_self - return { - "elapsed_min": round(elapsed / 60, 1), - "floor": round(floor, 4), - "shown_per_min": round(self.shown_per_min(), 2), - "counts": { - "total": self.total, - "skip_seen_self": self.skip_seen_self, - "scored": scored, - "fail_resolve": self.fail_resolve, - "fail_spam": self.fail_spam, - "shown": self.shown, - "windows": self.windows_elapsed, - "windows_empty": self.windows_empty, - }, - "rates": { - "shown_pct": round(100 * self.shown / scored, 4) if scored else 0, - "window_hit_pct": round(100 * (self.windows_elapsed - self.windows_empty) - / self.windows_elapsed, 1) if self.windows_elapsed else 0, - }, - "score_histogram": self._hist_to_dict(self.score_hist), - "shown_histogram": self._hist_to_dict(self.shown_hist), - "medoid_hits": dict(sorted(self.medoid_hits.items(), key=lambda x: -x[1])), - "recency": self.age_stats(), - } - - def write_file(self, floor): - data = self.to_dict(floor) - tmp = STATS_FILE + ".tmp" - with open(tmp, "w") as f: - json.dump(data, f, indent=2) - f.write("\n") - os.replace(tmp, STATS_FILE) - - def print_summary(self, floor): - d = self.to_dict(floor) - c = d["counts"] - r = d["rates"] - print(f"\n{'='*60}", file=sys.stderr) - print(f"Filter stats ({d['elapsed_min']} min, floor={d['floor']})", - file=sys.stderr) - print(f"{'='*60}", file=sys.stderr) - print(f" Total scanned: {c['total']:>8}", file=sys.stderr) - print(f" Skip seen/self: {c['skip_seen_self']:>8}", file=sys.stderr) - print(f" Scored: {c['scored']:>8}", file=sys.stderr) - print(f" Fail resolve: {c['fail_resolve']:>8}", file=sys.stderr) - print(f" Fail spam: {c['fail_spam']:>8}", file=sys.stderr) - print(f" Shown: {c['shown']:>8} " - f"({r['shown_pct']:.3f}% of scored)", file=sys.stderr) - print(f" Windows: {c['windows']:>8} " - f"({c['windows_empty']} empty, {r['window_hit_pct']:.0f}% hit)", - file=sys.stderr) - print(f"\n Shown/min (5m window): {d['shown_per_min']}", file=sys.stderr) - if d["score_histogram"]: - print(f"\n Score distribution (all scored posts):", file=sys.stderr) - for bucket, count in d["score_histogram"].items(): - bar = "#" * min(count, 60) - print(f" {bucket}: {count:>7} {bar}", file=sys.stderr) - if d["shown_histogram"]: - print(f"\n Shown post scores:", file=sys.stderr) - for bucket, count in d["shown_histogram"].items(): - bar = "#" * min(count, 60) - print(f" {bucket}: {count:>7} {bar}", file=sys.stderr) - if d["medoid_hits"]: - print(f"\n Medoid hits (shown posts):", file=sys.stderr) - for label, count in d["medoid_hits"].items(): - print(f" {count:>6} {label}", file=sys.stderr) - rec = d.get("recency") - if rec: - print(f"\n Post recency (5m window, {rec['count']} posts):", file=sys.stderr) - print(f" Median age: {rec['median_sec']:.0f}s", file=sys.stderr) - print(f" P90 age: {rec['p90_sec']:.0f}s", file=sys.stderr) - print(f" Range: {rec['min_sec']:.0f}s - {rec['max_sec']:.0f}s", file=sys.stderr) - print(f" Under 60s: {rec['pct_under_60s']:.0f}%", file=sys.stderr) - print(f" Under 5min: {rec['pct_under_300s']:.0f}%", file=sys.stderr) - print(f"\n Stats written to: {STATS_FILE}", file=sys.stderr) - print(f"{'='*60}", file=sys.stderr) - - -# ── Main ───────────────────────────────────────────────────────────────────── - -def normalize(v): - n = np.linalg.norm(v) - return v / n if n > 0 else v - - -def main(): - parser = argparse.ArgumentParser(description="Personalized Bluesky feed") - parser.add_argument("token", help="Divepool bearer token") - parser.add_argument("did", help="Your AT Protocol DID") - parser.add_argument("--window", type=int, required=True, - help="Selection window in seconds") - parser.add_argument("--top", type=int, default=3, - help="Max posts to show per window (default: 3)") - parser.add_argument("--floor", type=float, default=0.5, - help="Minimum composite score to consider (default: 0.5)") - args = parser.parse_args() - - token, user_did = args.token, args.did - floor = args.floor - - # ── Step 1: Get your post clusters + medoid embeddings ─────────────── - print("\nProfiling your posts...", file=sys.stderr, flush=True) - - medoids = [] - try: - data = divepool_search(token, "", limit=1000, did=user_did, cluster=True) - for cl in data.get("clusters", []): - emb = cl.get("medoid_embedding", []) - topics = cl.get("topics", []) - if emb and topics and cl.get("size", 0) >= 2: - label = " ".join(topics[:3]) - medoids.append((label, normalize(np.array(emb, dtype=np.float32)))) - except Exception as e: - print(f" error: {e}", file=sys.stderr, flush=True) - - print(f" {len(medoids)} clusters from your posts:", file=sys.stderr, flush=True) - for label, _ in medoids: - print(f" - {label}", file=sys.stderr, flush=True) - - if not medoids: - print("No interest clusters found.", file=sys.stderr) - return - - medoid_labels = [m[0] for m in medoids] - medoid_matrix = np.stack([m[1] for m in medoids]) # (N, dim) - print(f"\n{len(medoids)} medoids, window={args.window}s, top={args.top}, floor={floor}", - file=sys.stderr, flush=True) - - # ── Step 2: Stream firehose, score every post, pick best per window ── - print("Starting live feed... (Ctrl-C to stop)\n", file=sys.stderr, flush=True) - - firehose_buffer = [] - stop_event = threading.Event() - stream_thread = threading.Thread( - target=stream_firehose, - args=(token, firehose_buffer, stop_event), - daemon=True, - ) - stream_thread.start() - - seen_rkeys = set() - did_freq = defaultdict(int) - stats = FilterStats() - shown_scores = [] - - # Window state: candidates collected during current window - window_start = time.time() - # Each candidate: (composite_score, best_sim, best_label, other_labels, did, col, rkey) - candidates = [] - - def is_spam(did, text): - if did_freq[did] > 5: - return True - if len(text) < 20: - return True - if text.count("#") > 8: - return True - if text.count("http") > 3: - return True - return False - - def show(handle, text, rkey, comp_score, best_sim, label): - link = bsky_link(handle, rkey) - gold = "" - if len(shown_scores) >= 10: - mean = np.mean(shown_scores) - std = np.std(shown_scores) - if std > 0 and comp_score >= mean + 1.5 * std: - gold = " *" - shown_scores.append(comp_score) - print(f"[{comp_score:.2f}{gold} {label}]\n{text}\n{link} — @{handle}\n", flush=True) - stats.record_shown(best_sim) - - def flush_window(): - """Pick the best candidates from the window, resolve and show them.""" - nonlocal window_start, candidates - stats.windows_elapsed += 1 - - if not candidates: - stats.windows_empty += 1 - window_start = time.time() - candidates = [] - return - - # Sort by composite score, pick top N - candidates.sort(key=lambda c: -c[0]) - picks = candidates[:args.top] - - shown_in_window = 0 - for comp_score, best_sim, best_label, other_labels, did, col, rkey in picks: - handle, text = resolve_post_text(did, col, rkey) - if not handle or not text: - stats.fail_resolve += 1 - continue - if is_spam(did, text): - stats.fail_spam += 1 - continue - seen_rkeys.add(rkey) - if other_labels: - label = best_label + " + " + " + ".join(other_labels[:2]) - else: - label = best_label - stats.medoid_hits[best_label] += 1 - show(handle, text, rkey, comp_score, best_sim, label) - shown_in_window += 1 - - if shown_in_window == 0: - stats.windows_empty += 1 - - window_start = time.time() - candidates = [] - did_freq.clear() - - last_stats = time.time() - - try: - while True: - # Drain firehose buffer, score everything - while firehose_buffer: - did, col, rkey, lang, c_emb = firehose_buffer.pop(0) - stats.total += 1 - did_freq[did] += 1 - stats.record_age(rkey) - - if rkey in seen_rkeys or did == user_did: - stats.skip_seen_self += 1 - continue - - emb = normalize(np.array(c_emb, dtype=np.float32)) - sims = medoid_matrix @ emb # (N,) dot products - best_idx = int(np.argmax(sims)) - best_sim = float(sims[best_idx]) - - # Composite: best similarity + small bonus for multi-interest hits - # Multi-hit bar: must be genuinely similar, not just baseline noise - multi_bar = max(0.7, best_sim * 0.85) - multi_hits = [i for i in range(len(sims)) - if i != best_idx and sims[i] >= multi_bar] - comp_score = best_sim + min(0.15, 0.05 * len(multi_hits)) - stats.record_score(best_sim) # histogram tracks raw sim - - if comp_score < floor: - continue - - other_labels = [medoid_labels[i] for i in multi_hits[:2]] - candidates.append(( - comp_score, best_sim, medoid_labels[best_idx], - other_labels, did, col, rkey, - )) - # Keep candidate list bounded — only need top picks + some margin - # for resolve/spam failures. Sort and trim when it gets large. - max_keep = args.top * 5 - if len(candidates) > max_keep * 2: - candidates.sort(key=lambda c: -c[0]) - candidates = candidates[:max_keep] - - # Check if window is up - now = time.time() - if now - window_start >= args.window: - flush_window() - - # Write stats every 30s - if now - last_stats >= 30: - rate = stats.shown_per_min() - age = stats.age_stats() - age_str = f" age_med={age['median_sec']:.0f}s" if age else "" - n_cand = len(candidates) - best_cand = f" best={max(c[0] for c in candidates):.2f}" if candidates else "" - print( - f"[stats] scanned={stats.total} shown={stats.shown} " - f"rate={rate:.1f}/min cands={n_cand}{best_cand}{age_str}", - file=sys.stderr, flush=True, - ) - stats.write_file(floor) - last_stats = now - - time.sleep(0.05) - - except KeyboardInterrupt: - # Flush any remaining candidates - if candidates: - flush_window() - stats.write_file(floor) - stats.print_summary(floor) - stop_event.set() - - -if __name__ == "__main__": - main() diff --git a/experiments/for_you_v2.py b/experiments/for_you_v2.py deleted file mode 100644 index 1530dd1..0000000 --- a/experiments/for_you_v2.py +++ /dev/null @@ -1,735 +0,0 @@ -#!/usr/bin/env python3 -""" -For You v2 — personalized Bluesky feed from the Divepool firehose. -Improves on for_you.py: per-post embeddings instead of blurry medoid centroids, -IDF-style weighting to suppress generic interests. - -Up to 1000 actual post embeddings as reference vectors. Specificity weighting -ensures distinctive interests score higher than catch-all clusters. -Works with or without bearer token (128d vs 768d embeddings). - -Usage: - python3 for_you_v2.py [token] --window 60 -""" - -import argparse -import json -import os -import sys -import threading -import time -import urllib.error -import urllib.parse -import urllib.request -from collections import defaultdict - -import numpy as np - -DIVEPOOL_SEARCH = "https://divepool.social/api/v1/search" -DIVEPOOL_MEDOIDS = "https://divepool.social/api/v1/medoids" -DIVEPOOL_STREAM = "https://divepool.social/api/v1/embeddings" -BSKY_API = "https://public.api.bsky.app/xrpc" -STATS_FILE = os.path.join(os.path.dirname(os.path.abspath(__file__)), "for_you_v2_stats.json") -LOG_FILE = os.path.join(os.path.dirname(os.path.abspath(__file__)), "for_you_v2.log") - -# ── Scoring pipeline (applied in order) ──────────────────────────────────── -# -# 1. IDF-weighted similarity: raw_sims * ref_weights → best_sim -# Per-reference specificity weight suppresses generic catch-all matches. -# -# 2. Multi-interest bonus: comp_score = best_sim + bonus for cross-cluster hits -# Posts matching multiple user interests get a small boost (max +0.15). -# -# 3. Z-score normalization (per cluster, Welford's online): -# Normalizes comp_score relative to each cluster's running distribution, -# so a "good" match for a rare cluster competes fairly with common ones. -# Raw comp_score used as fallback during warmup (first 100 scored posts). -# -# --- above computed per firehose post; below computed per window at flush --- -# -# 4. Cluster diversity malus: -MALUS per prior win for that cluster. -# Prevents one dominant cluster from winning every window. -# -# 5. Account affinity bonus: +AFFINITY_BONUS * max_sim(account_medoids, user_centroids) -# Accounts whose overall posting profile overlaps the user's interests. -# -# 6. Account credibility: +CREDIBILITY_WEIGHT * credibility_score(-1 to +1) -# Penalizes spam bots, boosts small legit accounts, neutral for large ones. -# -WARMUP = 100 # min scored posts before z-scores kick in -MALUS = 0.5 # z-score penalty per prior cluster win -AFFINITY_BONUS = 1.0 # max z-score bonus for account affinity -CREDIBILITY_WEIGHT = 1.5 # scales credibility (-1 to +1) - - -# ── API helpers ────────────────────────────────────────────────────────────── - -def divepool_search(token, query, limit=100, did=None, cluster=False, - distinct=True, include_embeddings=False): - payload = {"query": query, "limit": limit, "distinct": distinct} - if did: - payload["did"] = did - if cluster: - payload["cluster"] = True - if include_embeddings: - payload["include_embeddings"] = True - data = json.dumps(payload).encode() - headers = {"Content-Type": "application/json"} - if token: - headers["Authorization"] = f"Bearer {token}" - req = urllib.request.Request(DIVEPOOL_SEARCH, data=data, headers=headers, - method="POST") - with urllib.request.urlopen(req, timeout=60) as resp: - return json.loads(resp.read()) - - -def fetch_account_medoids(token, dids): - """Batch-fetch up to 3 cluster medoids per account (max 25 DIDs).""" - payload = {"dids": dids[:25]} - data = json.dumps(payload).encode() - headers = {"Content-Type": "application/json"} - if token: - headers["Authorization"] = f"Bearer {token}" - req = urllib.request.Request(DIVEPOOL_MEDOIDS, data=data, headers=headers, - method="POST") - try: - with urllib.request.urlopen(req, timeout=10) as resp: - return json.loads(resp.read()) - except Exception: - return {"accounts": {}} - - -def fetch_bsky_profile(did): - """Fetch public profile stats. Returns (followers, following, posts) or None.""" - url = f"{BSKY_API}/app.bsky.actor.getProfile?actor={urllib.parse.quote(did)}" - req = urllib.request.Request(url, headers={"Accept": "application/json"}) - try: - with urllib.request.urlopen(req, timeout=5) as resp: - data = json.loads(resp.read()) - return ( - data.get("followersCount", 0), - data.get("followsCount", 0), - data.get("postsCount", 0), - ) - except Exception: - return None - - -def account_credibility(followers, following, posts): - """Score from -1 (spam) through 0 (neutral/big) to +1 (small & legit). - Penalizes spam bots; boosts approachable small accounts; neutral for large ones.""" - if posts == 0: - return -0.5 # zero-post accounts are almost always fake - - ratio = followers / posts - - # Base spam signal: followers/posts ratio - # ratio 0.08 → -0.6, ratio 0.3 → 0.0, ratio 0.75 → +0.5, ratio 1+ → +0.6 - base = (ratio - 0.3) / (ratio + 0.2) - - # Suspicious following patterns (follow-back bots, fake follower farms) - if followers > 0: - ff_ratio = following / followers - if ff_ratio > 10: - base = min(base, -0.5) - elif ff_ratio > 5: - base *= 0.5 - elif following > 100: - base = min(base, -0.5) - - # Small & legit bonus: accounts under ~200 followers with healthy ratios - # are the approachable long-tail we want to surface - if ratio >= 0.3 and followers < 200: - base += 0.3 # boost small legit accounts - - # Big accounts: cap positive score — they don't need help - if followers > 1000: - base = min(base, 0.1) - - return max(-1.0, min(1.0, base)) - - -def resolve_post_text(did, collection, rkey): - uri = f"at://{did}/{collection}/{rkey}" - url = f"{BSKY_API}/app.bsky.feed.getPostThread?uri={urllib.parse.quote(uri)}&depth=0" - req = urllib.request.Request(url, headers={"Accept": "application/json"}) - try: - with urllib.request.urlopen(req, timeout=5) as resp: - data = json.loads(resp.read()) - post = data.get("thread", {}).get("post", {}) - text = post.get("record", {}).get("text", "") - handle = post.get("author", {}).get("handle", "") - return handle, text - except Exception: - return None, None - - -def bsky_link(handle, rkey): - return f"https://bsky.app/profile/{handle}/post/{rkey}" - - -def tid_to_timestamp(tid): - """Decode an AT Protocol TID (base32-sortable) to a Unix timestamp in seconds.""" - charset = "234567abcdefghijklmnopqrstuvwxyz" - try: - n = 0 - for ch in tid: - n = n * 32 + charset.index(ch) - usec = n >> 10 - return usec / 1_000_000 - except (ValueError, IndexError): - return 0.0 - - -# ── Firehose streaming ────────────────────────────────────────────────────── - -def stream_firehose(token, post_queue, stop_event): - import zstandard as zstd - while not stop_event.is_set(): - try: - headers = {} - if token: - headers["Authorization"] = f"Bearer {token}" - req = urllib.request.Request(DIVEPOOL_STREAM, headers=headers) - with urllib.request.urlopen(req, timeout=30) as resp: - dctx = zstd.ZstdDecompressor() - reader = dctx.stream_reader(resp) - buf = b"" - while not stop_event.is_set(): - chunk = reader.read(8192) - if not chunk: - break - buf += chunk - while b"\n" in buf: - line, buf = buf.split(b"\n", 1) - if not line.strip(): - continue - try: - batch = json.loads(line) - except json.JSONDecodeError: - continue - if not isinstance(batch, dict): - continue - dids = batch.get("did") - if not isinstance(dids, list) or not dids: - continue - cols = batch.get("col", []) - rkeys = batch.get("rkey", []) - langs = batch.get("lang", []) - cs = batch.get("c", []) - for i in range(len(dids)): - col = cols[i] if i < len(cols) else "" - if "feed.post" not in col: - continue - c_emb = cs[i] if i < len(cs) else [] - if not c_emb: - continue - rkey = rkeys[i] if i < len(rkeys) else "" - lang = langs[i] if i < len(langs) else "" - post_queue.append((dids[i], col, rkey, lang, c_emb)) - except Exception: - if not stop_event.is_set(): - time.sleep(2) - - -# ── Stats tracker ────────────────────────────────────────────────────────── - -class FilterStats: - """Tracks scoring distribution and selection stats.""" - - BINS = [0.0, 0.3, 0.4, 0.5, 0.55, 0.6, 0.65, 0.7, 0.75, 0.8, 0.85, 0.9, 0.95, 1.0] - - def __init__(self): - self.started = time.time() - self.total = 0 - self.skip_seen_self = 0 - self.fail_resolve = 0 - self.fail_spam = 0 - self.shown = 0 - self.windows_elapsed = 0 - self.windows_empty = 0 - self.score_hist = np.zeros(len(self.BINS) - 1, dtype=np.int64) - self.shown_hist = np.zeros(len(self.BINS) - 1, dtype=np.int64) - self.cluster_hits = defaultdict(int) - self._shown_times = [] - self._recent_ages = [] - - def record_age(self, rkey): - ts = tid_to_timestamp(rkey) - if ts > 0: - age = time.time() - ts - self._recent_ages.append((time.time(), age)) - - def age_stats(self): - now = time.time() - cutoff = now - 300 - self._recent_ages = [(t, a) for t, a in self._recent_ages if t > cutoff] - if not self._recent_ages: - return None - ages = [a for _, a in self._recent_ages] - return { - "count": len(ages), - "median_sec": round(float(np.median(ages)), 1), - "p90_sec": round(float(np.percentile(ages, 90)), 1), - "max_sec": round(max(ages), 1), - "min_sec": round(min(ages), 1), - "pct_under_60s": round(100 * sum(1 for a in ages if a < 60) / len(ages), 1), - "pct_under_300s": round(100 * sum(1 for a in ages if a < 300) / len(ages), 1), - } - - def record_score(self, score): - idx = np.searchsorted(self.BINS, score, side="right") - 1 - idx = max(0, min(idx, len(self.score_hist) - 1)) - self.score_hist[idx] += 1 - - def record_shown(self, score): - self.shown += 1 - self._shown_times.append(time.time()) - idx = np.searchsorted(self.BINS, score, side="right") - 1 - idx = max(0, min(idx, len(self.shown_hist) - 1)) - self.shown_hist[idx] += 1 - - def shown_per_min(self): - now = time.time() - cutoff = now - 300 - self._shown_times = [t for t in self._shown_times if t > cutoff] - if not self._shown_times: - return 0.0 - window = now - self._shown_times[0] - if window < 10: - return 0.0 - return len(self._shown_times) / (window / 60) - - def _hist_to_dict(self, hist): - out = {} - for i, count in enumerate(hist): - if count > 0: - lo, hi = self.BINS[i], self.BINS[i + 1] - out[f"{lo:.2f}-{hi:.2f}"] = int(count) - return out - - def to_dict(self): - elapsed = time.time() - self.started - scored = self.total - self.skip_seen_self - return { - "elapsed_min": round(elapsed / 60, 1), - "shown_per_min": round(self.shown_per_min(), 2), - "counts": { - "total": self.total, - "skip_seen_self": self.skip_seen_self, - "scored": scored, - "fail_resolve": self.fail_resolve, - "fail_spam": self.fail_spam, - "shown": self.shown, - "windows": self.windows_elapsed, - "windows_empty": self.windows_empty, - }, - "rates": { - "shown_pct": round(100 * self.shown / scored, 4) if scored else 0, - "window_hit_pct": round(100 * (self.windows_elapsed - self.windows_empty) - / self.windows_elapsed, 1) if self.windows_elapsed else 0, - }, - "score_histogram": self._hist_to_dict(self.score_hist), - "shown_histogram": self._hist_to_dict(self.shown_hist), - "cluster_hits": dict(sorted(self.cluster_hits.items(), key=lambda x: -x[1])), - "recency": self.age_stats(), - } - - def write_file(self): - data = self.to_dict() - tmp = STATS_FILE + ".tmp" - with open(tmp, "w") as f: - json.dump(data, f, indent=2) - f.write("\n") - os.replace(tmp, STATS_FILE) - - def print_summary(self): - d = self.to_dict() - c = d["counts"] - r = d["rates"] - print(f"\n{'='*60}", file=sys.stderr) - print(f"Filter stats ({d['elapsed_min']} min)", - file=sys.stderr) - print(f"{'='*60}", file=sys.stderr) - print(f" Total scanned: {c['total']:>8}", file=sys.stderr) - print(f" Skip seen/self: {c['skip_seen_self']:>8}", file=sys.stderr) - print(f" Scored: {c['scored']:>8}", file=sys.stderr) - print(f" Fail resolve: {c['fail_resolve']:>8}", file=sys.stderr) - print(f" Fail spam: {c['fail_spam']:>8}", file=sys.stderr) - print(f" Shown: {c['shown']:>8} " - f"({r['shown_pct']:.3f}% of scored)", file=sys.stderr) - print(f" Windows: {c['windows']:>8} " - f"({c['windows_empty']} empty, {r['window_hit_pct']:.0f}% hit)", - file=sys.stderr) - print(f"\n Shown/min (5m window): {d['shown_per_min']}", file=sys.stderr) - def print_hist(title, hist): - if not hist: - return - peak = max(hist.values()) - print(f"\n {title}:", file=sys.stderr) - for bucket, count in hist.items(): - bar = "#" * max(1, round(50 * count / peak)) if count else "" - print(f" {bucket}: {count:>7} {bar}", file=sys.stderr) - - print_hist("Score distribution (all scored posts)", d["score_histogram"]) - print_hist("Shown post scores", d["shown_histogram"]) - if d["cluster_hits"]: - print(f"\n Cluster hits (shown posts):", file=sys.stderr) - for label, count in d["cluster_hits"].items(): - print(f" {count:>6} {label}", file=sys.stderr) - rec = d.get("recency") - if rec: - print(f"\n Post recency (5m window, {rec['count']} posts):", file=sys.stderr) - print(f" Median age: {rec['median_sec']:.0f}s", file=sys.stderr) - print(f" P90 age: {rec['p90_sec']:.0f}s", file=sys.stderr) - print(f" Range: {rec['min_sec']:.0f}s - {rec['max_sec']:.0f}s", file=sys.stderr) - print(f" Under 60s: {rec['pct_under_60s']:.0f}%", file=sys.stderr) - print(f" Under 5min: {rec['pct_under_300s']:.0f}%", file=sys.stderr) - print(f"\n Stats written to: {STATS_FILE}", file=sys.stderr) - print(f"{'='*60}", file=sys.stderr) - - -# ── Main ───────────────────────────────────────────────────────────────────── - -def normalize(v): - n = np.linalg.norm(v) - return v / n if n > 0 else v - - -def main(): - parser = argparse.ArgumentParser(description="Personalized Bluesky feed (v2 — per-post embeddings)") - parser.add_argument("token", nargs="?", default="", - help="Divepool bearer token (optional — without token, 128d embeddings)") - parser.add_argument("did", help="Your AT Protocol DID") - parser.add_argument("--window", type=int, required=True, - help="Selection window in seconds") - args = parser.parse_args() - - token, user_did = args.token, args.did - - # ── Step 1: Fetch your posts with per-result embeddings ───────────── - print("\nProfiling your posts...", file=sys.stderr, flush=True) - - cluster_id_to_label = {} - ref_labels = [] - ref_cluster_ids = [] - ref_vecs = [] - - try: - data = divepool_search(token, "", limit=1000, did=user_did, - cluster=True, include_embeddings=True) - - # Build cluster label lookup - for cl in data.get("clusters", []): - topics = cl.get("topics", []) - if topics: - cluster_id_to_label[cl["id"]] = " ".join(topics[:3]) - - # Extract per-result embeddings - for r in data.get("results", []): - emb = r.get("embedding", []) - if not emb: - continue - cid = r.get("cluster_id") - if cid is None or cid not in cluster_id_to_label: - continue # skip unclustered posts — too vague as reference - label = cluster_id_to_label[cid] - - ref_labels.append(label) - ref_cluster_ids.append(cid) - ref_vecs.append(normalize(np.array(emb, dtype=np.float32))) - - except Exception as e: - print(f" error: {e}", file=sys.stderr, flush=True) - - if not ref_vecs: - print("No reference embeddings found.", file=sys.stderr) - return - - ref_matrix = np.stack(ref_vecs) # (N, dim) - n_clusters = len(set(ref_cluster_ids) - {-1}) - - print(f" {len(ref_vecs)} reference embeddings from {n_clusters} clusters:", - file=sys.stderr, flush=True) - for cid, label in sorted(cluster_id_to_label.items()): - count = ref_cluster_ids.count(cid) - print(f" - {label} ({count} posts)", file=sys.stderr, flush=True) - - # ── Step 2: IDF-style weighting — downweight generic references ───── - print(" Computing specificity weights...", file=sys.stderr, flush=True) - self_sims = ref_matrix @ ref_matrix.T # (N, N) - np.fill_diagonal(self_sims, 0) - mean_sims = self_sims.mean(axis=1) # how similar each ref is to all others - ref_weights = 1.0 - mean_sims # distinctive refs get higher weight - ref_weights = ref_weights / ref_weights.max() # normalize: most distinctive = 1.0 - ref_weights = np.clip(ref_weights, 0.3, 1.0) # floor at 0.3 so nothing is crushed - - # Show weight distribution per cluster - for cid, label in sorted(cluster_id_to_label.items()): - mask = [i for i, c in enumerate(ref_cluster_ids) if c == cid] - if mask: - w = ref_weights[mask] - print(f" {label}: weight {w.mean():.2f} (min={w.min():.2f}, max={w.max():.2f})", - file=sys.stderr, flush=True) - - # ── User cluster centroids (for account affinity scoring) ──────────── - user_centroids = [] - for cid in sorted(cluster_id_to_label): - mask = [i for i, c in enumerate(ref_cluster_ids) if c == cid] - if mask: - centroid = ref_matrix[mask].mean(axis=0) - centroid = centroid / np.linalg.norm(centroid) - user_centroids.append(centroid) - user_centroid_matrix = np.stack(user_centroids) # (K, dim) - - print(f"\n{len(ref_vecs)} refs, {n_clusters} clusters, " - f"window={args.window}s", - file=sys.stderr, flush=True) - - # ── Step 3: Stream firehose, score every post, pick best per window ── - print("Starting live feed... (Ctrl-C to stop)\n", file=sys.stderr, flush=True) - - firehose_buffer = [] - stop_event = threading.Event() - stream_thread = threading.Thread( - target=stream_firehose, - args=(token, firehose_buffer, stop_event), - daemon=True, - ) - stream_thread.start() - - seen_rkeys = set() - did_freq = defaultdict(int) - stats = FilterStats() - - window_start = time.time() - candidates = [] # (z, comp_score, best_label, other_labels, did, col, rkey) - cluster_shown_count = defaultdict(int) # how many times each cluster won a window - account_affinity_cache = {} # did -> affinity score (0-1) - profile_cache = {} # did -> (followers, following, posts) - - # Per-cluster running stats for z-score normalization (Welford's online algo) - cluster_stats = defaultdict(lambda: [0, 0.0, 0.0]) # [count, mean, M2] - total_scored = 0 - - def update_cluster_stats(label, score): - s = cluster_stats[label] - s[0] += 1 - delta = score - s[1] - s[1] += delta / s[0] - s[2] += delta * (score - s[1]) - - def z_score(label, score): - s = cluster_stats[label] - if s[0] < 5: - return score # not enough data, fall back to raw - std = (s[2] / s[0]) ** 0.5 - if std < 1e-6: - return 0.0 - return (score - s[1]) / std - - def is_spam(did, text): - if did_freq[did] > 5: - return True - if len(text) < 20: - return True - if text.count("#") > 8: - return True - if text.count("http") > 3: - return True - return False - - def format_age(rkey): - age_sec = time.time() - tid_to_timestamp(rkey) - if age_sec < 90: - return f"{age_sec:.0f} seconds ago" - elif age_sec < 3600: - return f"{age_sec / 60:.0f} minutes ago" - elif age_sec < 86400: - return f"{age_sec / 3600:.1f} hours ago" - else: - return f"{age_sec / 86400:.1f} days ago" - - def show(handle, text, rkey, adj_z, comp_score, label, affinity=0.0, cred=0.0): - link = bsky_link(handle, rkey) - aff_str = f" aff={affinity:.2f}" if affinity > 0 else "" - cred_str = f" cred={cred:+.2f}" if abs(cred) > 0.1 else "" - print(f"[z={adj_z:.1f} {comp_score:.2f}{aff_str}{cred_str} {label}]\n" - f"{text}\n{link} — @{handle}\n" - f"{format_age(rkey)}\n", - flush=True) - stats.record_shown(comp_score) - - def flush_window(): - """Pick the best candidate by z-score, resolve and show.""" - nonlocal window_start, candidates - stats.windows_elapsed += 1 - - if not candidates: - stats.windows_empty += 1 - window_start = time.time() - candidates = [] - return - - # Fetch account medoids + profiles for top candidates - candidate_dids = list(dict.fromkeys( - d for _, _, _, _, d, _, _ in sorted(candidates, key=lambda c: -c[0])[:20] - )) - uncached_medoids = [d for d in candidate_dids[:25] if d not in account_affinity_cache] - if uncached_medoids: - resp = fetch_account_medoids(token, uncached_medoids) - for did_key, acct in resp.get("accounts", {}).items(): - medoids = acct.get("medoids", []) - if not medoids: - account_affinity_cache[did_key] = 0.0 - continue - med_vecs = [] - for m in medoids: - e = m.get("embedding", []) - if e: - med_vecs.append(normalize(np.array(e, dtype=np.float32))) - if not med_vecs: - account_affinity_cache[did_key] = 0.0 - continue - med_matrix = np.stack(med_vecs) # (up to 3, dim) - sims = med_matrix @ user_centroid_matrix.T # (3, K) - account_affinity_cache[did_key] = float(sims.max()) - for d in uncached_medoids: - if d not in account_affinity_cache: - account_affinity_cache[d] = 0.0 - - uncached_profiles = [d for d in candidate_dids[:10] if d not in profile_cache] - for d in uncached_profiles: - prof = fetch_bsky_profile(d) - profile_cache[d] = prof if prof else (0, 0, 0) - - # Apply cumulative malus + account affinity + credibility - adjusted = [] - for z, cs, bl, ol, d, c, rk in candidates: - affinity = account_affinity_cache.get(d, 0.0) - prof = profile_cache.get(d) - cred = account_credibility(*prof) if prof else 0.0 - adj = (z - - MALUS * cluster_shown_count[bl] - + AFFINITY_BONUS * affinity - + CREDIBILITY_WEIGHT * cred) - adjusted.append((adj, z, cs, bl, ol, d, c, rk, cred)) - adjusted.sort(key=lambda c: -c[0]) # sort by adjusted z-score - - shown = False - for adj_z, z, comp_score, best_label, other_labels, did, col, rkey, cred in adjusted: - handle, text = resolve_post_text(did, col, rkey) - if not handle or not text: - stats.fail_resolve += 1 - continue - if is_spam(did, text): - stats.fail_spam += 1 - continue - seen_rkeys.add(rkey) - if other_labels: - label = best_label + " + " + " + ".join(other_labels[:2]) - else: - label = best_label - stats.cluster_hits[best_label] += 1 - cluster_shown_count[best_label] += 1 - affinity = account_affinity_cache.get(did, 0.0) - show(handle, text, rkey, adj_z, comp_score, label, affinity, cred) - shown = True - break # only show top 1 - - if not shown: - stats.windows_empty += 1 - - # One-line summary after flush - rate = stats.shown_per_min() - print( - f"[window {stats.windows_elapsed}] scanned={stats.total} shown={stats.shown} " - f"rate={rate:.1f}/min", - file=sys.stderr, flush=True, - ) - stats.write_file() - - window_start = time.time() - candidates = [] - did_freq.clear() - - last_file_log = time.time() - - try: - while True: - # Drain firehose buffer, score everything - while firehose_buffer: - did, col, rkey, lang, c_emb = firehose_buffer.pop(0) - stats.total += 1 - did_freq[did] += 1 - stats.record_age(rkey) - - if rkey in seen_rkeys or did == user_did: - stats.skip_seen_self += 1 - continue - - emb = normalize(np.array(c_emb, dtype=np.float32)) - raw_sims = ref_matrix @ emb # (N,) - weighted_sims = raw_sims * ref_weights # IDF-weighted - - best_idx = int(np.argmax(weighted_sims)) - best_sim = float(weighted_sims[best_idx]) - - # Multi-interest: count distinct clusters with high hits - multi_bar = max(0.7, best_sim * 0.85) - best_cluster = ref_cluster_ids[best_idx] - hit_clusters = set() - for i in range(len(weighted_sims)): - if weighted_sims[i] >= multi_bar and ref_cluster_ids[i] != best_cluster: - hit_clusters.add(ref_cluster_ids[i]) - comp_score = best_sim + min(0.15, 0.05 * len(hit_clusters)) - stats.record_score(best_sim) - - best_label = ref_labels[best_idx] - update_cluster_stats(best_label, comp_score) - total_scored += 1 - - z = z_score(best_label, comp_score) if total_scored >= WARMUP else comp_score - - other_labels = [cluster_id_to_label.get(cid, "?") - for cid in list(hit_clusters)[:2]] - candidates.append(( - z, comp_score, best_label, - other_labels, did, col, rkey, - )) - # Keep candidate list bounded (generous: flush adjustments can reorder) - if len(candidates) > 50: - candidates.sort(key=lambda c: -c[0]) - candidates = candidates[:25] - - # Check if window is up - now = time.time() - if now - window_start >= args.window: - flush_window() - - # Periodic log to file (not terminal) - if now - last_file_log >= 30: - age = stats.age_stats() - age_str = f" age_med={age['median_sec']:.0f}s" if age else "" - n_cand = len(candidates) - best_cand = f" best={max(c[0] for c in candidates):.2f}" if candidates else "" - with open(LOG_FILE, "a") as lf: - lf.write( - f"[{time.strftime('%H:%M:%S')}] scanned={stats.total} " - f"shown={stats.shown} cands={n_cand}{best_cand}{age_str}\n" - ) - stats.write_file() - last_file_log = now - - time.sleep(0.05) - - except KeyboardInterrupt: - if candidates: - flush_window() - stats.write_file() - stats.print_summary() - stop_event.set() - - -if __name__ == "__main__": - main() diff --git a/experiments/live_trends.py b/experiments/live_trends.py deleted file mode 100644 index b6787b3..0000000 --- a/experiments/live_trends.py +++ /dev/null @@ -1,304 +0,0 @@ -#!/usr/bin/env python3 -""" -Live Trends — taps the Divepool embedding firehose for a window, -clusters the incoming embeddings in real time, then uses the search API -to label each cluster with human-readable topics. - -Requires: numpy, scikit-learn (for local HDBSCAN/UMAP-lite clustering) -Falls back to simpler k-means if HDBSCAN unavailable. - -Usage: - python3 live_trends.py [--seconds 30] [--top 5] -""" - -import argparse -import json -import math -import struct -import sys -import time -import urllib.request -from collections import Counter, defaultdict - -import numpy as np - -DIVEPOOL_STREAM = "https://divepool.social/api/v1/embeddings" -DIVEPOOL_SEARCH = "https://divepool.social/api/v1/search" - - -def stream_embeddings(token: str, seconds: int) -> list[dict]: - """Collect embeddings from the firehose for `seconds` seconds.""" - req = urllib.request.Request(DIVEPOOL_STREAM, headers={ - "Authorization": f"Bearer {token}", - }) - events = [] - print(f" Streaming for {seconds}s...", end="", flush=True) - t0 = time.time() - - try: - with urllib.request.urlopen(req, timeout=seconds + 10) as resp: - # The stream is zstd-compressed NDJSON — we need to decompress. - # Use zstandard if available, else fall back to subprocess. - try: - import zstandard as zstd - dctx = zstd.ZstdDecompressor() - reader = dctx.stream_reader(resp) - except ImportError: - # Fall back to piping through zstd CLI - import subprocess, io - proc = subprocess.Popen( - ["zstd", "-d"], - stdin=subprocess.PIPE, - stdout=subprocess.PIPE, - stderr=subprocess.DEVNULL, - ) - # We need to pump data from resp to proc.stdin in a thread - import threading - def pump(): - try: - while True: - chunk = resp.read(8192) - if not chunk: - break - proc.stdin.write(chunk) - except Exception: - pass - finally: - proc.stdin.close() - t = threading.Thread(target=pump, daemon=True) - t.start() - reader = proc.stdout - - buf = b"" - while time.time() - t0 < seconds: - chunk = reader.read(4096) - if not chunk: - break - buf += chunk - while b"\n" in buf: - line, buf = buf.split(b"\n", 1) - if not line.strip(): - continue - try: - batch = json.loads(line) - except json.JSONDecodeError: - continue - # Batch is columnar: parallel arrays keyed by "did", "col", etc. - # Skip heartbeats (empty batches) or malformed data. - if not isinstance(batch, dict): - continue - dids = batch.get("did") - if not isinstance(dids, list) or len(dids) == 0: - continue - cols = batch.get("col", []) - rkeys = batch.get("rkey", []) - langs = batch.get("lang", []) - cs = batch.get("c", []) - rs = batch.get("r", []) - for i in range(len(dids)): - c_emb = cs[i] if i < len(cs) else [] - r_emb = rs[i] if i < len(rs) else [] - col = cols[i] if i < len(cols) else "" - lang = langs[i] if i < len(langs) else "" - rkey = rkeys[i] if i < len(rkeys) else "" - if c_emb: - events.append({ - "did": dids[i], - "col": col, - "rkey": rkey, - "lang": lang, - "c": c_emb, - "r": r_emb, - }) - if len(events) % 100 == 0: - print(f"\r Streaming for {seconds}s... {len(events)} events", end="", flush=True) - except Exception as e: - print(f"\n Stream ended: {e}") - - elapsed = time.time() - t0 - print(f"\r Collected {len(events)} events in {elapsed:.1f}s ({len(events)/max(elapsed,1):.0f}/s)") - return events - - -def cluster_embeddings(events: list[dict], n_clusters: int = 8) -> list[dict]: - """Cluster events by their C (cluster) embeddings using k-means.""" - if len(events) < n_clusters * 2: - n_clusters = max(2, len(events) // 5) - - from sklearn.cluster import MiniBatchKMeans - from sklearn.preprocessing import normalize - - # Build embedding matrix - C = np.array([e["c"] for e in events], dtype=np.float32) - C = normalize(C) # L2 normalize - - km = MiniBatchKMeans(n_clusters=n_clusters, random_state=42, n_init=3) - labels = km.fit_predict(C) - - # Group events by cluster - clusters = defaultdict(list) - for i, label in enumerate(labels): - clusters[label].append(events[i]) - - # Sort clusters by size - result = [] - for label, members in sorted(clusters.items(), key=lambda x: -len(x[1])): - # Compute centroid - centroid = np.mean([m["c"] for m in members], axis=0) - result.append({ - "id": label, - "size": len(members), - "members": members, - "centroid": centroid.tolist(), - "langs": Counter(m["lang"] for m in members).most_common(3), - "collections": Counter(m["col"] for m in members).most_common(3), - }) - return result - - -def label_cluster_via_search(token: str, cluster: dict) -> dict | None: - """Use a representative member's embedding to find similar indexed posts.""" - # Pick a few members and search for their content to get labels - members = cluster["members"] - # Use a random sample member's AT URI to find its text via search - # Instead, we'll search with the cluster centroid's nearest indexed neighbor - # by searching for posts from the same DIDs - sample = members[:3] - sample_dids = [m["did"] for m in sample] - - # We can't search by embedding directly, but we can search for content - # that's semantically similar. Let's use the collection type as a hint. - col_type = cluster["collections"][0][0] if cluster["collections"] else "" - - # Search for posts by these authors — the search API takes text queries, - # so we'll use a representative DID handle lookup - # Actually, let's just use the search API with clustering to find what topics - # match this cluster's centroid. We'll pick a member and look up its text. - return None # We'll use a different approach below - - -def search_for_context(token: str, query: str, limit: int = 50) -> dict: - payload = json.dumps({ - "query": query, "limit": limit, - "cluster": True, "distinct": True, - }).encode() - req = urllib.request.Request(DIVEPOOL_SEARCH, data=payload, headers={ - "Content-Type": "application/json", - "Authorization": f"Bearer {token}", - }, method="POST") - with urllib.request.urlopen(req, timeout=15) as resp: - return json.loads(resp.read()) - - -def resolve_posts(events: list[dict], max_posts: int = 5) -> list[str]: - """Resolve AT URIs to post text via the Bluesky API.""" - texts = [] - for e in events[:max_posts]: - uri = f"at://{e['did']}/{e['col']}/{e['rkey']}" - url = f"https://public.api.bsky.app/xrpc/app.bsky.feed.getPostThread?uri={urllib.parse.quote(uri)}&depth=0" - req = urllib.request.Request(url, headers={"Accept": "application/json"}) - try: - with urllib.request.urlopen(req, timeout=5) as resp: - data = json.loads(resp.read()) - post = data.get("thread", {}).get("post", {}).get("record", {}) - text = post.get("text", "") - if text: - texts.append(text) - except Exception: - pass - return texts - - -import urllib.parse - - -def main(): - parser = argparse.ArgumentParser(description="Live Bluesky trend detector") - parser.add_argument("token", help="Divepool bearer token") - parser.add_argument("--seconds", type=int, default=30, help="Streaming window (default: 30)") - parser.add_argument("--clusters", type=int, default=8, help="Number of clusters (default: 8)") - parser.add_argument("--resolve", action="store_true", help="Resolve post text via Bluesky API (slower)") - args = parser.parse_args() - - print(f"Live Trends — Bluesky Embedding Firehose") - print(f"{'─'*50}") - - # Step 1: Stream - print(f"\n[1/3] Tapping firehose...") - events = stream_embeddings(args.token, args.seconds) - if len(events) < 10: - print("Too few events to cluster. Try a longer window.") - return - - # Filter to posts only - posts = [e for e in events if "feed.post" in e.get("col", "")] - profiles = [e for e in events if "actor.profile" in e.get("col", "")] - print(f" {len(posts)} posts, {len(profiles)} profile updates") - - if len(posts) < 10: - print("Too few posts to cluster.") - return - - # Step 2: Cluster - print(f"\n[2/3] Clustering {len(posts)} post embeddings...") - clusters = cluster_embeddings(posts, n_clusters=args.clusters) - print(f" {len(clusters)} clusters formed") - - # Step 3: Label clusters by resolving sample posts - print(f"\n[3/3] Resolving cluster content...") - for cl in clusters: - members = cl["members"] - # Resolve a few posts from each cluster to understand content - if args.resolve: - texts = resolve_posts(members, max_posts=5) - else: - texts = [] - cl["sample_texts"] = texts - - # Language breakdown - lang_str = ", ".join(f"{lang}({n})" for lang, n in cl["langs"]) - cl["lang_str"] = lang_str - - # Unique authors - cl["unique_authors"] = len(set(m["did"] for m in members)) - - # ── Report ── - print(f"\n\n{'━'*50}") - print(f" LIVE TRENDS — {len(posts)} posts in {args.seconds}s") - print(f" ({len(posts)/args.seconds:.1f} posts/sec)") - print(f"{'━'*50}") - - # Overall language distribution - all_langs = Counter(e["lang"] for e in posts) - print(f"\n Language mix: {', '.join(f'{l}({n})' for l, n in all_langs.most_common(10))}") - - # Collection breakdown - all_cols = Counter(e["col"] for e in events) - print(f" Content types: {', '.join(f'{c.split('.')[-1]}({n})' for c, n in all_cols.most_common(5))}") - - for i, cl in enumerate(clusters): - pct = cl["size"] / len(posts) * 100 - print(f"\n ── Cluster #{cl['id']} — {cl['size']} posts ({pct:.0f}%) — {cl['unique_authors']} authors ──") - print(f" Languages: {cl['lang_str']}") - - if cl["sample_texts"]: - print(f" Sample posts:") - for t in cl["sample_texts"][:3]: - text = t[:120].replace("\n", " ") - print(f" \"{text}\"") - else: - # Show AT URIs for manual inspection - print(f" Sample URIs (use --resolve to fetch text):") - for m in cl["members"][:3]: - print(f" at://{m['did']}/{m['col']}/{m['rkey']}") - - # Embedding dimensionality check - if posts: - dim = len(posts[0].get("c", [])) - print(f"\n Embedding dimensionality: {dim}d (768d = authenticated)") - - print(f"\n{'━'*50}\n") - - -if __name__ == "__main__": - main() diff --git a/experiments/outliers.py b/experiments/outliers.py deleted file mode 100644 index 8ddb4b7..0000000 --- a/experiments/outliers.py +++ /dev/null @@ -1,366 +0,0 @@ -#!/usr/bin/env python3 -""" -Outlier Finder — taps the Divepool firehose, then uses UMAP + HDBSCAN to find -the surprising, niche, and weird corners of Bluesky. - -Unlike k-means, HDBSCAN: - - Finds clusters of varying density and size (micro-communities surface) - - Labels points that don't fit anywhere as noise (-1) — true outliers - - Provides per-point outlier scores for ranking weirdness - -We invert the usual presentation: smallest clusters and highest-outlier posts -come first. The big boring blobs get a one-line summary at the bottom. - -Usage: - python3 outliers.py [--seconds 600] [--min-cluster 3] -""" - -import argparse -import json -import sys -import time -import urllib.error -import urllib.parse -import urllib.request -from collections import Counter, defaultdict -from concurrent.futures import ThreadPoolExecutor, as_completed - -import numpy as np - -DIVEPOOL_STREAM = "https://divepool.social/api/v1/embeddings" -BSKY_API = "https://public.api.bsky.app/xrpc" - - -# ── Firehose streaming ────────────────────────────────────────────────────── - -def stream_embeddings(token: str, seconds: int) -> list[dict]: - req = urllib.request.Request(DIVEPOOL_STREAM, headers={ - "Authorization": f"Bearer {token}", - }) - events = [] - print(f" Streaming for {seconds}s...", end="", flush=True) - t0 = time.time() - - try: - import zstandard as zstd - with urllib.request.urlopen(req, timeout=seconds + 10) as resp: - dctx = zstd.ZstdDecompressor() - reader = dctx.stream_reader(resp) - buf = b"" - while time.time() - t0 < seconds: - chunk = reader.read(8192) - if not chunk: - break - buf += chunk - while b"\n" in buf: - line, buf = buf.split(b"\n", 1) - if not line.strip(): - continue - try: - batch = json.loads(line) - except json.JSONDecodeError: - continue - if not isinstance(batch, dict): - continue - dids = batch.get("did") - if not isinstance(dids, list) or len(dids) == 0: - continue - cols = batch.get("col", []) - rkeys = batch.get("rkey", []) - langs = batch.get("lang", []) - cs = batch.get("c", []) - rs = batch.get("r", []) - for i in range(len(dids)): - c_emb = cs[i] if i < len(cs) else [] - if not c_emb: - continue - col = cols[i] if i < len(cols) else "" - lang = langs[i] if i < len(langs) else "" - rkey = rkeys[i] if i < len(rkeys) else "" - r_emb = rs[i] if i < len(rs) else [] - events.append({ - "did": dids[i], "col": col, "rkey": rkey, - "lang": lang, "c": c_emb, "r": r_emb, - }) - if len(events) % 200 == 0 and len(events) > 0: - print(f"\r Streaming for {seconds}s... {len(events)} events", end="", flush=True) - except Exception as e: - print(f"\n Stream ended: {e}") - - elapsed = time.time() - t0 - print(f"\r Collected {len(events)} events in {elapsed:.1f}s ({len(events)/max(elapsed,1):.0f}/s)") - return events - - -# ── Clustering ─────────────────────────────────────────────────────────────── - -def cluster_with_hdbscan(events: list[dict], min_cluster_size: int = 5): - """ - UMAP (768d → 15d) then HDBSCAN. Returns (labels, outlier_scores, umap_2d). - """ - import umap - import hdbscan - - C = np.array([e["c"] for e in events], dtype=np.float32) - # L2 normalize - norms = np.linalg.norm(C, axis=1, keepdims=True) - norms[norms == 0] = 1 - C = C / norms - - print(f" UMAP 768d → 15d...", end="", flush=True) - t0 = time.time() - reducer = umap.UMAP( - n_components=15, n_neighbors=30, min_dist=0.0, - metric="cosine", random_state=42, low_memory=True, - ) - embedding_15d = reducer.fit_transform(C) - print(f" ({time.time()-t0:.1f}s)") - - print(f" HDBSCAN (min_cluster_size={min_cluster_size})...", end="", flush=True) - t0 = time.time() - clusterer = hdbscan.HDBSCAN( - min_cluster_size=min_cluster_size, - min_samples=2, - cluster_selection_method="eom", # excess of mass — favors many small clusters - prediction_data=True, - ) - labels = clusterer.fit_predict(embedding_15d) - outlier_scores = clusterer.outlier_scores_ - print(f" ({time.time()-t0:.1f}s)") - - # Also get 2D for potential visualization - print(f" UMAP 15d → 2d...", end="", flush=True) - t0 = time.time() - reducer_2d = umap.UMAP( - n_components=2, n_neighbors=30, min_dist=0.1, - metric="euclidean", random_state=42, low_memory=True, - ) - embedding_2d = reducer_2d.fit_transform(embedding_15d) - print(f" ({time.time()-t0:.1f}s)") - - return labels, outlier_scores, embedding_2d - - -# ── Post resolution ────────────────────────────────────────────────────────── - -def resolve_post(event: dict) -> str | None: - uri = f"at://{event['did']}/{event['col']}/{event['rkey']}" - url = f"{BSKY_API}/app.bsky.feed.getPostThread?uri={urllib.parse.quote(uri)}&depth=0" - req = urllib.request.Request(url, headers={"Accept": "application/json"}) - try: - with urllib.request.urlopen(req, timeout=5) as resp: - data = json.loads(resp.read()) - post = data.get("thread", {}).get("post", {}) - text = post.get("record", {}).get("text", "") - handle = post.get("author", {}).get("handle", "") - likes = post.get("likeCount", 0) - return f"@{handle} [{likes}♥] \"{text}\"" if text else None - except Exception: - return None - - -def resolve_batch(events: list[dict], max_posts: int = 5) -> list[str]: - results = [] - with ThreadPoolExecutor(max_workers=10) as pool: - futures = {pool.submit(resolve_post, e): e for e in events[:max_posts]} - for f in as_completed(futures): - r = f.result() - if r: - results.append(r) - return results - - -# ── Main ───────────────────────────────────────────────────────────────────── - -def main(): - parser = argparse.ArgumentParser(description="Bluesky outlier/niche finder") - parser.add_argument("token", help="Divepool bearer token") - parser.add_argument("--seconds", type=int, default=600, help="Streaming window (default: 600)") - parser.add_argument("--min-cluster", type=int, default=3, help="HDBSCAN min_cluster_size (default: 3, lower = more micro-clusters)") - args = parser.parse_args() - - print(f"Outlier Finder — Bluesky Firehose") - print(f"{'─'*55}") - - # ── Stream ─────────────────────────────────────────────────────────── - print(f"\n[1/4] Tapping firehose...") - events = stream_embeddings(args.token, args.seconds) - - posts = [e for e in events if "feed.post" in e.get("col", "")] - profiles = [e for e in events if "actor.profile" in e.get("col", "")] - print(f" {len(posts)} posts, {len(profiles)} profile updates") - - if len(posts) < 20: - print("Too few posts. Try a longer window.") - return - - # ── Cluster ────────────────────────────────────────────────────────── - print(f"\n[2/4] UMAP + HDBSCAN clustering...") - labels, outlier_scores, coords_2d = cluster_with_hdbscan(posts, args.min_cluster) - - # Attach labels and scores to events - for i, e in enumerate(posts): - e["cluster"] = int(labels[i]) - e["outlier_score"] = float(outlier_scores[i]) - e["x"] = float(coords_2d[i, 0]) - e["y"] = float(coords_2d[i, 1]) - - # Build cluster groups - clusters = defaultdict(list) - noise = [] - for e in posts: - if e["cluster"] == -1: - noise.append(e) - else: - clusters[e["cluster"]].append(e) - - n_clusters = len(clusters) - sizes = sorted([len(m) for m in clusters.values()]) - print(f" {n_clusters} clusters, {len(noise)} noise points (true outliers)") - if sizes: - print(f" Cluster sizes: min={sizes[0]}, median={sizes[len(sizes)//2]}, max={sizes[-1]}") - - # ── Resolve posts ──────────────────────────────────────────────────── - # We want to resolve: all small clusters + top outliers + a sample of big clusters. - small_clusters = {k: v for k, v in clusters.items() if len(v) <= 15} - big_clusters = {k: v for k, v in clusters.items() if len(v) > 15} - - # Count how many posts we need to resolve - to_resolve = [] - for members in small_clusters.values(): - to_resolve.extend(members[:5]) - noise_ranked = sorted(noise, key=lambda e: -e["outlier_score"]) - to_resolve.extend(noise_ranked[:20]) - for members in big_clusters.values(): - to_resolve.extend(members[:2]) - - # Deduplicate by rkey - seen = set() - unique_resolve = [] - for e in to_resolve: - if e["rkey"] not in seen: - seen.add(e["rkey"]) - unique_resolve.append(e) - - print(f"\n[3/4] Resolving {len(unique_resolve)} posts via Bluesky API...") - t0 = time.time() - resolved = {} - with ThreadPoolExecutor(max_workers=15) as pool: - futures = {pool.submit(resolve_post, e): e for e in unique_resolve} - for f in as_completed(futures): - e = futures[f] - r = f.result() - if r: - resolved[e["rkey"]] = r - print(f" Resolved {len(resolved)}/{len(unique_resolve)} ({time.time()-t0:.1f}s)") - - # ── Save 2D coordinates for optional visualization ─────────────────── - print(f"\n[4/4] Saving 2D map to outlier_map.json...") - map_data = [] - for e in posts: - map_data.append({ - "x": e["x"], "y": e["y"], - "cluster": e["cluster"], - "outlier_score": e["outlier_score"], - "did": e["did"], "rkey": e["rkey"], - "lang": e["lang"], - "text": resolved.get(e["rkey"], ""), - }) - with open("outlier_map.json", "w") as f: - json.dump(map_data, f) - print(f" {len(map_data)} points saved") - - # ── REPORT ─────────────────────────────────────────────────────────── - print(f"\n\n{'━'*55}") - print(f" OUTLIER REPORT — {len(posts)} posts, {n_clusters} clusters, {len(noise)} outliers") - print(f"{'━'*55}") - - # ── Section 1: True outliers (noise points, ranked by outlier score) - print(f"\n{'─'*55}") - print(f" TRUE OUTLIERS — posts that fit nowhere") - print(f" (HDBSCAN noise points, ranked by outlier score)") - print(f"{'─'*55}") - shown = 0 - for e in noise_ranked: - text = resolved.get(e["rkey"]) - if not text: - continue - # Truncate for display - lines = text.split('"') - display = text[:200].replace("\n", " ") - print(f"\n [{e['outlier_score']:.3f}] {display}") - shown += 1 - if shown >= 15: - break - - # ── Section 2: Micro-clusters (the niche communities) - micro = {k: v for k, v in clusters.items() if len(v) <= 10} - small = {k: v for k, v in clusters.items() if 10 < len(v) <= 30} - big = {k: v for k, v in clusters.items() if len(v) > 30} - - if micro: - print(f"\n{'─'*55}") - print(f" MICRO-CLUSTERS — tiny niche communities ({len(micro)} found)") - print(f"{'─'*55}") - for cid, members in sorted(micro.items(), key=lambda x: len(x[1])): - langs = Counter(m["lang"] for m in members).most_common(2) - lang_str = ", ".join(f"{l}" for l, _ in langs) - n_authors = len(set(m["did"] for m in members)) - print(f"\n Cluster #{cid} — {len(members)} posts, {n_authors} authors [{lang_str}]") - for m in members[:4]: - text = resolved.get(m["rkey"]) - if text: - display = text[:180].replace("\n", " ") - print(f" {display}") - - if small: - print(f"\n{'─'*55}") - print(f" SMALL CLUSTERS — emerging topics ({len(small)} found)") - print(f"{'─'*55}") - for cid, members in sorted(small.items(), key=lambda x: len(x[1])): - langs = Counter(m["lang"] for m in members).most_common(2) - lang_str = ", ".join(f"{l}" for l, _ in langs) - n_authors = len(set(m["did"] for m in members)) - print(f"\n Cluster #{cid} — {len(members)} posts, {n_authors} authors [{lang_str}]") - for m in members[:3]: - text = resolved.get(m["rkey"]) - if text: - display = text[:180].replace("\n", " ") - print(f" {display}") - - # ── Section 3: Big clusters (one-liner summary) - if big: - print(f"\n{'─'*55}") - print(f" BIG CLUSTERS — the expected mainstream ({len(big)} found)") - print(f"{'─'*55}") - for cid, members in sorted(big.items(), key=lambda x: -len(x[1])): - pct = len(members) / len(posts) * 100 - langs = Counter(m["lang"] for m in members).most_common(1) - n_authors = len(set(m["did"] for m in members)) - sample = resolved.get(members[0]["rkey"], "") - snippet = sample[:100].replace("\n", " ") if sample else "(unresolved)" - print(f" #{cid:3d} {len(members):4d} posts ({pct:4.1f}%) {n_authors:3d} authors {snippet}") - - # ── Language outliers - all_langs = Counter(e["lang"] for e in posts) - rare_langs = [(l, n) for l, n in all_langs.most_common() if n <= 5 and l] - if rare_langs: - print(f"\n{'─'*55}") - print(f" RARE LANGUAGE POSTS") - print(f"{'─'*55}") - for lang, count in rare_langs: - lang_posts = [e for e in posts if e["lang"] == lang] - print(f"\n {lang} ({count} posts):") - for e in lang_posts[:3]: - text = resolved.get(e["rkey"]) - if text: - display = text[:180].replace("\n", " ") - print(f" {display}") - - print(f"\n{'━'*55}") - print(f" 2D map saved to outlier_map.json ({len(posts)} points)") - print(f"{'━'*55}\n") - - -if __name__ == "__main__": - main() diff --git a/experiments/requirements.txt b/experiments/requirements.txt deleted file mode 100644 index f107310..0000000 --- a/experiments/requirements.txt +++ /dev/null @@ -1,5 +0,0 @@ -numpy>=1.26.0 -zstandard>=0.20.0 -scikit-learn>=1.3.0 -umap-learn>=0.5.0 -hdbscan>=0.8.33 diff --git a/experiments/topic_explorer.py b/experiments/topic_explorer.py deleted file mode 100644 index af528d5..0000000 --- a/experiments/topic_explorer.py +++ /dev/null @@ -1,288 +0,0 @@ -#!/usr/bin/env python3 -""" -Topic Explorer — combines Divepool semantic search with the public Bluesky API -to map the community structure behind any topic. - -For a given query it: -1. Searches Divepool for semantically relevant posts (with clustering) -2. Enriches results with Bluesky engagement data (likes, reposts, replies) -3. Resolves author profiles (follower counts, bios) -4. Checks follow relationships between top authors to find community clusters -5. Finds related posts by top authors to see what else they talk about -6. Outputs a rich topic report -""" - -import json -import sys -import time -import urllib.request -import urllib.error -import urllib.parse -from collections import defaultdict -from concurrent.futures import ThreadPoolExecutor, as_completed - -DIVEPOOL = "https://divepool.social/api/v1/search" -BSKY_API = "https://public.api.bsky.app/xrpc" - - -# ── Divepool search ────────────────────────────────────────────────────────── - -def divepool_search(token: str, query: str, limit: int = 300) -> dict: - payload = json.dumps({ - "query": query, "limit": limit, - "cluster": True, "distinct": True, - }).encode() - req = urllib.request.Request(DIVEPOOL, data=payload, headers={ - "Content-Type": "application/json", - "Authorization": f"Bearer {token}", - }, method="POST") - with urllib.request.urlopen(req, timeout=30) as resp: - return json.loads(resp.read()) - - -# ── Bluesky public API helpers ─────────────────────────────────────────────── - -def bsky_get(method: str, params: dict) -> dict | None: - qs = urllib.parse.urlencode(params) - url = f"{BSKY_API}/{method}?{qs}" - req = urllib.request.Request(url, headers={"Accept": "application/json"}) - try: - with urllib.request.urlopen(req, timeout=10) as resp: - return json.loads(resp.read()) - except urllib.error.HTTPError: - return None - - -def get_profiles(dids: list[str]) -> dict[str, dict]: - """Batch-fetch profiles (up to 25 per call).""" - profiles = {} - for i in range(0, len(dids), 25): - batch = dids[i:i+25] - qs = "&".join(f"actors={urllib.parse.quote(d)}" for d in batch) - url = f"{BSKY_API}/app.bsky.actor.getProfiles?{qs}" - req = urllib.request.Request(url, headers={"Accept": "application/json"}) - try: - with urllib.request.urlopen(req, timeout=10) as resp: - data = json.loads(resp.read()) - for p in data.get("profiles", []): - profiles[p["did"]] = p - except urllib.error.HTTPError: - pass - return profiles - - -def get_post_thread(uri: str) -> dict | None: - return bsky_get("app.bsky.feed.getPostThread", {"uri": uri, "depth": 0}) - - -def get_follows(did: str, limit: int = 100) -> list[str]: - """Get DIDs that `did` follows.""" - data = bsky_get("app.bsky.graph.getFollows", {"actor": did, "limit": limit}) - if not data: - return [] - return [f["did"] for f in data.get("follows", [])] - - -# ── Main logic ─────────────────────────────────────────────────────────────── - -def main(): - if len(sys.argv) < 3: - print("Usage: python3 topic_explorer.py ") - print('Example: python3 topic_explorer.py TOKEN "climate activism"') - sys.exit(1) - - token = sys.argv[1] - query = " ".join(sys.argv[2:]) - - print(f"Topic Explorer: \"{query}\"") - print(f"{'─'*60}") - - # ── Step 1: Semantic search ────────────────────────────────────────── - print("\n[1/5] Semantic search via Divepool...") - t0 = time.time() - search_data = divepool_search(token, query) - results = search_data.get("results", []) - clusters = search_data.get("clusters", []) - print(f" {len(results)} results, {len(clusters)} clusters ({time.time()-t0:.1f}s)") - - if not results: - print("No results found.") - return - - # ── Step 2: Enrich with engagement data ────────────────────────────── - print("\n[2/5] Fetching engagement data from Bluesky...") - t0 = time.time() - top_results = results[:30] # enrich top 30 - engagement = {} - - def fetch_engagement(r): - uri = f"at://{r['did']}/{r['collection']}/{r['rkey']}" - thread = get_post_thread(uri) - if thread and "thread" in thread: - post = thread["thread"].get("post", {}) - return r["rkey"], { - "likes": post.get("likeCount", 0), - "reposts": post.get("repostCount", 0), - "replies": post.get("replyCount", 0), - "uri": uri, - } - return r["rkey"], None - - with ThreadPoolExecutor(max_workers=10) as pool: - futures = [pool.submit(fetch_engagement, r) for r in top_results] - for f in as_completed(futures): - rkey, data = f.result() - if data: - engagement[rkey] = data - - print(f" Enriched {len(engagement)}/{len(top_results)} posts ({time.time()-t0:.1f}s)") - - # ── Step 3: Resolve author profiles ────────────────────────────────── - print("\n[3/5] Resolving author profiles...") - t0 = time.time() - unique_dids = list(dict.fromkeys(r["did"] for r in top_results))[:25] - profiles = get_profiles(unique_dids) - print(f" {len(profiles)} profiles resolved ({time.time()-t0:.1f}s)") - - # ── Step 4: Check follow relationships between top authors ─────────── - print("\n[4/5] Mapping follow graph between top authors...") - t0 = time.time() - top_dids = unique_dids[:15] # check top 15 - follow_graph: dict[str, set[str]] = {} - top_did_set = set(top_dids) - - def fetch_follows(did): - follows = get_follows(did) - mutual = set(follows) & top_did_set - return did, mutual - - with ThreadPoolExecutor(max_workers=8) as pool: - futures = [pool.submit(fetch_follows, did) for did in top_dids] - for f in as_completed(futures): - did, mutual = f.result() - if mutual - {did}: - follow_graph[did] = mutual - {did} - - total_edges = sum(len(v) for v in follow_graph.values()) - print(f" {total_edges} follow links among top {len(top_dids)} authors ({time.time()-t0:.1f}s)") - - # ── Step 5: Second-order search — what else do top authors post about? - print("\n[5/5] Probing what top authors talk about (second-order search)...") - t0 = time.time() - # Pick 3 contrasting queries based on cluster topics - alt_queries = [] - for cl in sorted(clusters, key=lambda c: c["size"], reverse=True)[:3]: - topics = cl.get("topics", []) - if topics: - alt_queries.append(topics[0]) - - alt_results = {} - for aq in alt_queries: - try: - d = divepool_search(token, aq, limit=100) - alt_results[aq] = d.get("results", []) - except Exception: - pass - print(f" Ran {len(alt_queries)} sub-queries ({time.time()-t0:.1f}s)") - - # ── REPORT ─────────────────────────────────────────────────────────── - - def handle_for(did): - p = profiles.get(did, {}) - return p.get("handle", did.split(":")[-1]) - - print(f"\n\n{'━'*60}") - print(f" TOPIC REPORT: \"{query}\"") - print(f"{'━'*60}") - - # Cluster overview - print(f"\n── Topic Clusters ──") - for cl in sorted(clusters, key=lambda c: c["size"], reverse=True): - topics = ", ".join(cl.get("topics", [])[:5]) - print(f"\n Cluster #{cl['id']} ({cl['size']} posts)") - print(f" Topics: {topics}") - # Sample posts from this cluster with engagement - shown = 0 - for idx in cl.get("result_indices", [])[:3]: - if idx < len(results): - r = results[idx] - eng = engagement.get(r["rkey"]) - eng_str = "" - if eng: - eng_str = f" [{eng['likes']}♥ {eng['reposts']}⟳ {eng['replies']}💬]" - text = r["text"][:100].replace("\n", " ") - print(f" @{r.get('handle', '?'):25s}{eng_str}") - print(f" \"{text}...\"") - shown += 1 - - # Top voices with profile context - print(f"\n── Top Voices ──") - scored = [] - for r in top_results: - eng = engagement.get(r["rkey"], {}) - impact = eng.get("likes", 0) + eng.get("reposts", 0) * 2 - scored.append((r, impact)) - scored.sort(key=lambda x: (-x[1], x[0]["score"])) - - for r, impact in scored[:10]: - p = profiles.get(r["did"], {}) - handle = p.get("handle", r.get("handle", "?")) - followers = p.get("followersCount", 0) - bio = (p.get("description") or "")[:80].replace("\n", " ") - eng = engagement.get(r["rkey"], {}) - eng_str = "" - if eng: - eng_str = f"{eng['likes']}♥ {eng['reposts']}⟳ {eng['replies']}💬" - - print(f"\n @{handle}") - print(f" {followers:,} followers | similarity={r['score']:.3f} | {eng_str}") - if bio: - print(f" Bio: \"{bio}\"") - text = r["text"][:120].replace("\n", " ") - print(f" Post: \"{text}...\"") - - # Follow network - if follow_graph: - print(f"\n── Community Network (who follows whom among top authors) ──") - # Find clusters of mutual follows - for did, follows in sorted(follow_graph.items(), key=lambda x: -len(x[1])): - src = handle_for(did) - targets = ", ".join(f"@{handle_for(d)}" for d in follows) - print(f" @{src} → {targets}") - - # Identify the most-followed within the topic - in_degree: dict[str, int] = defaultdict(int) - for did, follows in follow_graph.items(): - for f in follows: - in_degree[f] += 1 - if in_degree: - print(f"\n Hub accounts (most followed within topic):") - for did, count in sorted(in_degree.items(), key=lambda x: -x[1])[:5]: - print(f" @{handle_for(did)} — followed by {count}/{len(top_dids)} top authors") - - # Cross-topic presence - if alt_results: - print(f"\n── Related Conversations ──") - top_handles_main = {r.get("handle") for r in top_results} - for aq, ares in alt_results.items(): - alt_handles = {r.get("handle") for r in ares} - overlap = top_handles_main & alt_handles - {None, ""} - print(f"\n Sub-topic: \"{aq}\"") - if overlap: - print(f" Overlapping voices: {', '.join('@'+h for h in list(overlap)[:5])}") - else: - print(f" No overlapping voices (distinct sub-community)") - # Show top post from this sub-topic - if ares: - r = ares[0] - text = r["text"][:120].replace("\n", " ") - print(f" Top hit: @{r.get('handle', '?')}: \"{text}...\"") - - print(f"\n{'━'*60}") - print(f" Done. {len(results)} posts → {len(clusters)} clusters → " - f"{len(profiles)} profiles → {total_edges} follow links") - print(f"{'━'*60}\n") - - -if __name__ == "__main__": - main() diff --git a/experiments/vibe_map.py b/experiments/vibe_map.py deleted file mode 100644 index 05c00f9..0000000 --- a/experiments/vibe_map.py +++ /dev/null @@ -1,138 +0,0 @@ -#!/usr/bin/env python3 -""" -Bluesky Vibe Map — semantic pulse of the network via Divepool search. - -Runs a batch of diverse queries against the Divepool search API with -clustering enabled, then prints a compact digest: top clusters per query, -trending handles, and cross-query topic overlaps. -""" - -import json -import sys -import time -import urllib.request - -BASE_URL = "https://divepool.social/api/v1/search" - -QUERIES = [ - "climate change action protest", - "open source software community", - "mental health support", - "indie game development", - "astronomy astrophotography space", - "cooking recipes food", - "live music concerts festival", - "labor union strike workers", - "generative AI art ethics", - "queer joy pride community", -] - -def search(token: str, query: str, limit: int = 200) -> dict: - payload = json.dumps({ - "query": query, - "limit": limit, - "cluster": True, - "distinct": True, - }).encode() - req = urllib.request.Request( - BASE_URL, - data=payload, - headers={ - "Content-Type": "application/json", - "Authorization": f"Bearer {token}", - }, - method="POST", - ) - with urllib.request.urlopen(req, timeout=30) as resp: - return json.loads(resp.read()) - - -def print_query_digest(query: str, data: dict, elapsed: float): - results = data.get("results", []) - clusters = data.get("clusters", []) - - print(f"\n{'='*60}") - print(f" \"{query}\"") - print(f" {len(results)} results, {len(clusters)} clusters, {elapsed:.1f}s") - print(f"{'='*60}") - - # Show top clusters sorted by size - for cl in sorted(clusters, key=lambda c: c["size"], reverse=True)[:5]: - topics = ", ".join(cl.get("topics", [])[:4]) - # Grab a sample post from this cluster - sample_idx = cl["result_indices"][0] if cl["result_indices"] else None - sample_text = "" - if sample_idx is not None and sample_idx < len(results): - sample_text = results[sample_idx]["text"][:120].replace("\n", " ") - print(f"\n Cluster #{cl['id']} ({cl['size']} posts)") - print(f" Topics: {topics}") - if sample_text: - print(f" Sample: \"{sample_text}...\"") - - # Top handles by score - top = sorted(results[:20], key=lambda r: r["score"])[:5] - print(f"\n Top handles:") - for r in top: - print(f" @{r['handle']:30s} score={r['score']:.3f}") - - -def main(): - if len(sys.argv) < 2: - print("Usage: python3 vibe_map.py [query ...]") - sys.exit(1) - - token = sys.argv[1] - queries = sys.argv[2:] if len(sys.argv) > 2 else QUERIES - - all_topics: dict[str, list[str]] = {} # topic -> list of queries it appeared in - all_handles: dict[str, int] = {} # handle -> count across queries - - print("Bluesky Vibe Map") - print(f"Scanning {len(queries)} queries...\n") - - for query in queries: - t0 = time.time() - try: - data = search(token, query) - except Exception as e: - print(f"\n ERROR on \"{query}\": {e}") - continue - elapsed = time.time() - t0 - - print_query_digest(query, data, elapsed) - - # Accumulate cross-query stats - for cl in data.get("clusters", []): - for topic in cl.get("topics", []): - all_topics.setdefault(topic, []).append(query) - for r in data.get("results", []): - h = r.get("handle", "") - if h: - all_handles[h] = all_handles.get(h, 0) + 1 - - # Cross-query summary - print(f"\n\n{'#'*60}") - print(f" CROSS-QUERY SUMMARY") - print(f"{'#'*60}") - - # Topics appearing across multiple queries - shared = {t: qs for t, qs in all_topics.items() if len(qs) > 1} - if shared: - print(f"\n Topics spanning multiple queries:") - for topic, qs in sorted(shared.items(), key=lambda x: -len(x[1]))[:10]: - print(f" \"{topic}\" — appears in {len(qs)} queries") - else: - print(f"\n No topics shared across queries (clusters are well-separated)") - - # Handles appearing in multiple queries (cross-topic posters) - multi = {h: c for h, c in all_handles.items() if c > 1} - if multi: - print(f"\n Cross-topic voices (appear in 2+ queries):") - for h, c in sorted(multi.items(), key=lambda x: -x[1])[:15]: - print(f" @{h} — {c} queries") - - print() - - -if __name__ == "__main__": - main() diff --git a/go.mod b/go.mod deleted file mode 100644 index 9d7b393..0000000 --- a/go.mod +++ /dev/null @@ -1,5 +0,0 @@ -module github.com/divepool/embedding_firehose_client - -go 1.25.0 - -require github.com/klauspost/compress v1.18.5 diff --git a/go.sum b/go.sum deleted file mode 100644 index 1c48397..0000000 --- a/go.sum +++ /dev/null @@ -1,2 +0,0 @@ -github.com/klauspost/compress v1.18.5 h1:/h1gH5Ce+VWNLSWqPzOVn6XBO+vJbCNGvjoaGBFW2IE= -github.com/klauspost/compress v1.18.5/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= diff --git a/main.go b/main.go deleted file mode 100644 index 1f0f66c..0000000 --- a/main.go +++ /dev/null @@ -1,353 +0,0 @@ -// Example client for the Divepool embedding firehose and search API (experimental). -// Both the firehose and this client are highly experimental and subject to change. -// -// Subcommands: -// -// go run . stream # public firehose (128d) -// go run . stream -token x -out events.jsonl # token firehose (768d) → file -// go run . search "query text" # public search (128d embeddings) -// go run . search -token x "query text" # bearer search (768d embeddings) -// go run . search -token x -cluster "query text" # search with clustering + topics -// go run . search -did did:plc:abc123 # browse account (clustered) -// go run . search -did did:plc:abc123 "query text" # search within account -// go run . search -did did:plc:abc123 -rkey 3abc # find similar posts in the account -// go run . search -langs de,en -since 2026-06-01T00:00:00Z "query" # language/time-filtered search -// go run . search -min-score 0.55 "query text" # drop weak matches -// go run . search -embeddings "query text" # include per-result embeddings -// go run . search -dids did:plc:a,did:plc:b,did:plc:c # cluster a set of accounts by their medoids -// go run . search -dids-file followers.txt # same, but read DIDs (one per line) from file -// go run . search -dids-file followers.txt -viewer-did did:plc:abc # …scoped to the languages that account posts in -// go run . medoids did:plc:abc123 did:plc:def456 # top 3 medoids per account (128d) -// go run . medoids -token x did:plc:abc123 # bearer medoids (768d) -// go run . similar did:plc:abc123 # nearest accounts (auto language scope) -// go run . similar -min-cluster-size 5 -langs de did:plc:abc123 # denoised, German topic matches only -// -// Without a subcommand, defaults to "stream" for backwards compatibility. -package main - -import ( - "bufio" - "encoding/json" - "flag" - "fmt" - "log" - "math" - "net/http" - "os" - "strings" - "time" - - "github.com/klauspost/compress/zstd" -) - -// heartbeatTimeout is 2x the server's 10s heartbeat interval. -// If no frame (data or heartbeat) arrives within this window, the connection is stale. -const heartbeatTimeout = 20 * time.Second - -// event is the per-item JSONL output format when -out is specified. -type event struct { - DID string `json:"did"` - Col string `json:"col"` - Rkey string `json:"rkey"` - Lang string `json:"lang"` - C []float32 `json:"c"` - R []float32 `json:"r"` -} - -func main() { - // Detect subcommand. Default to "stream" for backwards compatibility. - sub := "stream" - if len(os.Args) > 1 && !strings.HasPrefix(os.Args[1], "-") { - sub = os.Args[1] - os.Args = append(os.Args[:1], os.Args[2:]...) - } - - switch sub { - case "stream": - runStream() - case "search": - runSearch() - case "medoids": - runMedoids() - case "similar": - runSimilar() - default: - fmt.Fprintf(os.Stderr, "unknown subcommand: %s\nusage: %s [stream|search|medoids|similar] [flags]\n", sub, os.Args[0]) - os.Exit(1) - } -} - -func runStream() { - urlFlag := flag.String("url", "https://divepool.social", "base URL") - tokenFlag := flag.String("token", "", "bearer token") - outFlag := flag.String("out", "", "write JSONL events to file (append)") - flag.Parse() - - var outFile *os.File - if *outFlag != "" { - f, err := os.OpenFile(*outFlag, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0644) - if err != nil { - log.Fatalf("open %s: %v", *outFlag, err) - } - defer f.Close() - outFile = f - log.Printf("writing events to %s", *outFlag) - } - - for { - if err := stream(*urlFlag, *tokenFlag, outFile); err != nil { - log.Printf("stream error: %v, reconnecting in 5s", err) - time.Sleep(5 * time.Second) - } - } -} - -func runSearch() { - urlFlag := flag.String("url", "https://divepool.social", "base URL") - tokenFlag := flag.String("token", "", "bearer token (optional — non-bearer gets 128d embeddings)") - didFlag := flag.String("did", "", "scope to a specific account (DID)") - didsFlag := flag.String("dids", "", "cluster a set of accounts by their medoids (comma-separated DIDs, max 500)") - didsFileFlag := flag.String("dids-file", "", "read DIDs (one per line) for -dids mode") - viewerDIDFlag := flag.String("viewer-did", "", "restrict -dids clustering to the languages this account posts in (response echoes the derived langs)") - rkeyFlag := flag.String("rkey", "", "use this post's embedding as query (requires -did)") - limitFlag := flag.Int("limit", 20, "max results (server max: 400 global, 1200 DID-scoped)") - clusterFlag := flag.Bool("cluster", false, "enable clustering + topic extraction") - distinctFlag := flag.Bool("distinct", true, "one result per account") - embedFlag := flag.Bool("embeddings", false, "include per-result cluster embeddings") - langsFlag := flag.String("langs", "", "only posts in these languages (comma-separated ISO 639-1, e.g. de,en)") - sinceFlag := flag.String("since", "", "only posts created at/after this RFC3339 timestamp") - untilFlag := flag.String("until", "", "only posts created before this RFC3339 timestamp") - minScoreFlag := flag.Float64("min-score", math.NaN(), "minimum cosine similarity to the query (-1..1)") - flag.Parse() - - dids := parseDIDs(*didsFlag, *didsFileFlag) - - query := strings.Join(flag.Args(), " ") - if query != "" && *rkeyFlag != "" { - log.Fatal("query and -rkey are mutually exclusive") - } - if *rkeyFlag != "" && *didFlag == "" { - log.Fatal("-rkey requires -did") - } - if len(dids) > 0 && (*didFlag != "" || *rkeyFlag != "" || query != "") { - log.Fatal("-dids/-dids-file cannot be combined with -did, -rkey, or a query") - } - if query == "" && *didFlag == "" && *rkeyFlag == "" && len(dids) == 0 { - log.Fatal("need at least one of: query text, -did, -dids/-dids-file, -rkey") - } - if *viewerDIDFlag != "" && len(dids) == 0 { - log.Fatal("-viewer-did requires -dids/-dids-file") - } - - var langs []string - if *langsFlag != "" { - langs = strings.Split(*langsFlag, ",") - } - var minScore *float64 // NaN sentinel = flag not set → field omitted - if !math.IsNaN(*minScoreFlag) { - minScore = minScoreFlag - } - - resp, dur, err := Search(*urlFlag, *tokenFlag, SearchRequest{ - Query: query, - DID: *didFlag, - DIDs: dids, - ViewerDID: *viewerDIDFlag, - Rkey: *rkeyFlag, - Limit: *limitFlag, - Distinct: distinctFlag, - Cluster: *clusterFlag, - IncludeEmbeddings: *embedFlag, - Langs: langs, - Since: *sinceFlag, - Until: *untilFlag, - MinScore: minScore, - }) - if err != nil { - log.Fatalf("search failed: %v", err) - } - - fmt.Fprintf(os.Stderr, "%d results in %s\n", len(resp.Results), dur.Round(time.Millisecond)) - - enc := json.NewEncoder(os.Stdout) - enc.SetIndent("", " ") - if err := enc.Encode(resp); err != nil { - log.Fatalf("encode response: %v", err) - } -} - -// parseDIDs merges -dids (comma-separated) and -dids-file (one per line) into -// a deduplicated, trimmed slice. Blank lines and # comments in the file are skipped. -func parseDIDs(csv, file string) []string { - seen := map[string]bool{} - var out []string - add := func(s string) { - s = strings.TrimSpace(s) - if s == "" || strings.HasPrefix(s, "#") || seen[s] { - return - } - seen[s] = true - out = append(out, s) - } - for d := range strings.SplitSeq(csv, ",") { - add(d) - } - if file != "" { - data, err := os.ReadFile(file) - if err != nil { - log.Fatalf("read -dids-file %s: %v", file, err) - } - for line := range strings.SplitSeq(string(data), "\n") { - add(line) - } - } - return out -} - -func runSimilar() { - urlFlag := flag.String("url", "https://divepool.social", "base URL") - tokenFlag := flag.String("token", "", "bearer token (optional — output is identical, attribution only)") - limitFlag := flag.Int("limit", 30, "max accounts (server max 100)") - minClusterFlag := flag.Int("min-cluster-size", 0, "suppress topic clusters smaller than this, on both the account and its matches") - langsFlag := flag.String("langs", "", "language scope: comma-separated ISO 639-1, or 'all' for no filter (default: the account's own posting languages)") - flag.Parse() - - if flag.NArg() != 1 { - log.Fatal("need exactly one DID argument") - } - - // nil = server derives the subject's languages; &[] = no filter. - var langs *[]string - switch *langsFlag { - case "": - case "all": - langs = &[]string{} - default: - l := strings.Split(*langsFlag, ",") - langs = &l - } - - resp, dur, err := SimilarAccounts(*urlFlag, *tokenFlag, SimilarAccountsRequest{ - DID: flag.Arg(0), - Limit: *limitFlag, - MinClusterSize: *minClusterFlag, - Langs: langs, - }) - if err != nil { - log.Fatalf("similar accounts failed: %v", err) - } - - fmt.Fprintf(os.Stderr, "%d accounts in %s\n", len(resp.Accounts), dur.Round(time.Millisecond)) - - enc := json.NewEncoder(os.Stdout) - enc.SetIndent("", " ") - if err := enc.Encode(resp); err != nil { - log.Fatalf("encode response: %v", err) - } -} - -func runMedoids() { - urlFlag := flag.String("url", "https://divepool.social", "base URL") - tokenFlag := flag.String("token", "", "bearer token (optional — non-bearer gets 128d embeddings)") - flag.Parse() - - dids := flag.Args() - if len(dids) == 0 { - log.Fatal("need at least one DID argument") - } - if len(dids) > 25 { - log.Fatal("max 25 DIDs") - } - - resp, dur, err := Medoids(*urlFlag, *tokenFlag, MedoidsRequest{ - DIDs: dids, - }) - if err != nil { - log.Fatalf("medoids failed: %v", err) - } - - fmt.Fprintf(os.Stderr, "%d accounts in %s\n", len(resp.Accounts), dur.Round(time.Millisecond)) - - enc := json.NewEncoder(os.Stdout) - enc.SetIndent("", " ") - if err := enc.Encode(resp); err != nil { - log.Fatalf("encode response: %v", err) - } -} - -func stream(baseURL, token string, outFile *os.File) error { - req, err := http.NewRequest("GET", baseURL+"/api/v1/embeddings", nil) - if err != nil { - return fmt.Errorf("build request: %w", err) - } - if token != "" { - req.Header.Set("Authorization", "Bearer "+token) - } - resp, err := http.DefaultClient.Do(req) - if err != nil { - return fmt.Errorf("GET /embeddings: %w", err) - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - return fmt.Errorf("unexpected status: %d", resp.StatusCode) - } - - reader, err := zstd.NewReader(resp.Body) - if err != nil { - return fmt.Errorf("zstd reader: %w", err) - } - defer reader.Close() - - scanner := bufio.NewScanner(reader) - scanner.Buffer(make([]byte, 16<<20), 16<<20) // 16 MB — batches can be large - - // lines is fed by the scanner goroutine; closed on scanner EOF/error. - type scanResult struct { - line []byte - err error - } - lines := make(chan scanResult) - go func() { - defer close(lines) - for scanner.Scan() { - cp := make([]byte, len(scanner.Bytes())) - copy(cp, scanner.Bytes()) - lines <- scanResult{line: cp} - } - lines <- scanResult{err: scanner.Err()} - }() - - for { - select { - case res, ok := <-lines: - if !ok { - return fmt.Errorf("stream closed") - } - if res.err != nil { - return res.err - } - var b Batch - if err := json.Unmarshal(res.line, &b); err != nil { - continue - } - if b.Len() == 0 { - continue // heartbeat — resets the timeout - } - for i := range b.DID { - fmt.Printf("%s lang=%s cluster=%dd retrieval=%dd\n", - b.ATURI(i), b.Lang[i], len(b.C[i]), len(b.R[i])) - if outFile != nil { - row, _ := json.Marshal(event{ - DID: b.DID[i], Col: b.Col[i], Rkey: b.Rkey[i], - Lang: b.Lang[i], C: b.C[i], R: b.R[i], - }) - row = append(row, '\n') - outFile.Write(row) - } - } - case <-time.After(heartbeatTimeout): - return fmt.Errorf("no data received in %v (stale connection)", heartbeatTimeout) - } - } -} diff --git a/medoids.go b/medoids.go deleted file mode 100644 index da4a267..0000000 --- a/medoids.go +++ /dev/null @@ -1,73 +0,0 @@ -package main - -import ( - "bytes" - "encoding/json" - "fmt" - "io" - "net/http" - "strings" - "time" -) - -// MedoidsRequest is the POST /medoids request body. -type MedoidsRequest struct { - DIDs []string `json:"dids"` -} - -// MedoidsResponse is the POST /medoids response. -type MedoidsResponse struct { - Accounts map[string]AccountMedoids `json:"accounts"` -} - -// AccountMedoids holds medoids for a single account. -type AccountMedoids struct { - Medoids []Medoid `json:"medoids"` - AccountEmbedding []float32 `json:"account_embedding,omitempty"` // cluster-size-weighted medoid mean, L2-normalized (128d public, 768d bearer) -} - -// Medoid is a single cluster medoid. -type Medoid struct { - ClusterID int `json:"cluster_id"` - IsPrimary bool `json:"is_primary"` - Collection string `json:"collection"` - Rkey string `json:"rkey"` - ClusterSize int `json:"cluster_size"` - Embedding []float32 `json:"embedding"` // 128d public, 768d bearer - Topics []string `json:"topics,omitempty"` // c-TF-IDF topic keywords (en/de only) -} - -// Medoids calls POST /medoids with the given request. -func Medoids(baseURL, token string, req MedoidsRequest) (*MedoidsResponse, time.Duration, error) { - body, err := json.Marshal(req) - if err != nil { - return nil, 0, fmt.Errorf("marshal request: %w", err) - } - - httpReq, err := http.NewRequest("POST", baseURL+"/api/v1/medoids", bytes.NewReader(body)) - if err != nil { - return nil, 0, fmt.Errorf("build request: %w", err) - } - httpReq.Header.Set("Content-Type", "application/json") - if token != "" { - httpReq.Header.Set("Authorization", "Bearer "+token) - } - - start := time.Now() - resp, err := http.DefaultClient.Do(httpReq) - if err != nil { - return nil, 0, fmt.Errorf("POST /medoids: %w", err) - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - respBody, _ := io.ReadAll(resp.Body) - return nil, 0, fmt.Errorf("status %d: %s", resp.StatusCode, strings.TrimSpace(string(respBody))) - } - - var result MedoidsResponse - if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { - return nil, 0, fmt.Errorf("decode response: %w", err) - } - return &result, time.Since(start), nil -} diff --git a/requirements.txt b/requirements.txt deleted file mode 100644 index c7c43cd..0000000 --- a/requirements.txt +++ /dev/null @@ -1,6 +0,0 @@ -transformers>=4.45.0 -torch>=2.1.0 -numpy>=1.26.0 -safetensors>=0.4.0 -huggingface-hub>=0.20.0 -lingua-language-detector>=2.1.0 diff --git a/search.go b/search.go deleted file mode 100644 index 12fffaf..0000000 --- a/search.go +++ /dev/null @@ -1,100 +0,0 @@ -package main - -import ( - "bytes" - "encoding/json" - "fmt" - "io" - "net/http" - "strings" - "time" -) - -// SearchRequest is the POST /search request body. -type SearchRequest struct { - Query string `json:"query,omitempty"` // semantic search text (exclusive with rkey) - DID string `json:"did,omitempty"` // scope to a specific account (DID) - DIDs []string `json:"dids,omitempty"` // cluster a set of accounts by their top medoids (max 500; mutually exclusive with did/rkey/query) - ViewerDID string `json:"viewer_did,omitempty"` // dids mode only: restrict medoids to the languages this account posts in (derived server-side from its medoids) - Rkey string `json:"rkey,omitempty"` // use this post's embedding as query (requires did, exclusive with query) - Limit int `json:"limit,omitempty"` // default/max: 400 global, 1200 DID-scoped - Distinct *bool `json:"distinct,omitempty"` // default true (one post per account) - Cluster bool `json:"cluster,omitempty"` // enable UMAP+HDBSCAN clustering + topics (forced true for did-only and dids modes) - IncludeEmbeddings bool `json:"include_embeddings,omitempty"` // include per-result cluster embeddings - Langs []string `json:"langs,omitempty"` // only posts in these languages (any mode; mutually exclusive with viewer_did) - Since string `json:"since,omitempty"` // RFC3339: only posts created at/after (post modes; not dids) - Until string `json:"until,omitempty"` // RFC3339: only posts created before (post modes; not dids) - MinScore *float64 `json:"min_score,omitempty"` // minimum cosine similarity to the query (= -score); requires query or rkey -} - -// SearchResponse is the POST /search response. -type SearchResponse struct { - Results []SearchResult `json:"results"` - Clusters []SearchCluster `json:"clusters,omitempty"` - Langs []string `json:"langs,omitempty"` // applied language filter: explicit langs or viewer_did-derived (empty = none) -} - -// SearchResult is a single search hit. -type SearchResult struct { - DID string `json:"did"` - Handle string `json:"handle,omitempty"` - Collection string `json:"collection"` - Rkey string `json:"rkey"` - Text string `json:"text"` - Score float64 `json:"score"` - CreatedAt string `json:"created_at,omitempty"` - DetectedLang string `json:"detected_lang,omitempty"` - ClusterID *int `json:"cluster_id,omitempty"` - Topics []string `json:"topics,omitempty"` - Embedding []float32 `json:"embedding,omitempty"` // cluster embedding (128d non-bearer, 768d bearer) -} - -// SearchCluster describes a topic cluster when cluster=true. -type SearchCluster struct { - ID int `json:"id"` - Size int `json:"size"` - Topics []string `json:"topics,omitempty"` - ResultIndices []int `json:"result_indices"` - MedoidIndex *int `json:"medoid_index,omitempty"` // index into results of most representative post - MedoidEmbedding []float32 `json:"medoid_embedding,omitempty"` // medoid embedding (128d non-bearer, 768d bearer) -} - -// ATURI returns the AT URI for the result. -func (r *SearchResult) ATURI() string { - return fmt.Sprintf("at://%s/%s/%s", r.DID, r.Collection, r.Rkey) -} - -// Search calls POST /search with the given request. -func Search(baseURL, token string, req SearchRequest) (*SearchResponse, time.Duration, error) { - body, err := json.Marshal(req) - if err != nil { - return nil, 0, fmt.Errorf("marshal request: %w", err) - } - - httpReq, err := http.NewRequest("POST", baseURL+"/api/v1/search", bytes.NewReader(body)) - if err != nil { - return nil, 0, fmt.Errorf("build request: %w", err) - } - httpReq.Header.Set("Content-Type", "application/json") - if token != "" { - httpReq.Header.Set("Authorization", "Bearer "+token) - } - - start := time.Now() - resp, err := http.DefaultClient.Do(httpReq) - if err != nil { - return nil, 0, fmt.Errorf("POST /search: %w", err) - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - respBody, _ := io.ReadAll(resp.Body) - return nil, 0, fmt.Errorf("status %d: %s", resp.StatusCode, strings.TrimSpace(string(respBody))) - } - - var result SearchResponse - if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { - return nil, 0, fmt.Errorf("decode response: %w", err) - } - return &result, time.Since(start), nil -} diff --git a/similar.go b/similar.go deleted file mode 100644 index 5a962dd..0000000 --- a/similar.go +++ /dev/null @@ -1,82 +0,0 @@ -package main - -import ( - "bytes" - "encoding/json" - "fmt" - "io" - "net/http" - "strings" - "time" -) - -// SimilarAccountsRequest is the POST /similar-accounts request body. -type SimilarAccountsRequest struct { - DID string `json:"did"` - Limit int `json:"limit,omitempty"` // default/max 100; margin baseline always uses the full candidate pool - MinClusterSize int `json:"min_cluster_size,omitempty"` // suppress one-off topic clusters on both subject and candidate side - // Langs distinguishes absent from empty: nil = auto-scope to the subject's - // posting languages, pointer to empty = no language filter, explicit list - // applies as-is. - Langs *[]string `json:"langs,omitempty"` -} - -// SimilarAccountsResponse is the POST /similar-accounts response. -type SimilarAccountsResponse struct { - Accounts []SimilarAccount `json:"accounts"` - ScoreMean float64 `json:"score_mean"` // candidate-pool mean: score = margin + score_mean - Langs []string `json:"langs,omitempty"` // applied language scope (explicit or subject-derived) -} - -// SimilarAccount is one ranked similar account. -type SimilarAccount struct { - DID string `json:"did"` - Handle string `json:"handle,omitempty"` - Score float64 `json:"score"` // cross-topic LSE affinity — relative, compare within one response - Margin float64 `json:"margin"` // score − score_mean; the discriminative number, ≤0 = noise - // Matched is the account's medoid post nearest to one of the subject's - // medoids — fetch did+collection+rkey via the Bluesky API to render why. - Matched MatchedMedoid `json:"matched"` -} - -// MatchedMedoid identifies the candidate's best-matching topic cluster. -type MatchedMedoid struct { - Collection string `json:"collection"` - Rkey string `json:"rkey"` - ClusterSize int `json:"cluster_size"` -} - -// SimilarAccounts calls POST /similar-accounts with the given request. -func SimilarAccounts(baseURL, token string, req SimilarAccountsRequest) (*SimilarAccountsResponse, time.Duration, error) { - body, err := json.Marshal(req) - if err != nil { - return nil, 0, fmt.Errorf("marshal request: %w", err) - } - - httpReq, err := http.NewRequest("POST", baseURL+"/api/v1/similar-accounts", bytes.NewReader(body)) - if err != nil { - return nil, 0, fmt.Errorf("build request: %w", err) - } - httpReq.Header.Set("Content-Type", "application/json") - if token != "" { - httpReq.Header.Set("Authorization", "Bearer "+token) - } - - start := time.Now() - resp, err := http.DefaultClient.Do(httpReq) - if err != nil { - return nil, 0, fmt.Errorf("POST /similar-accounts: %w", err) - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - respBody, _ := io.ReadAll(resp.Body) - return nil, 0, fmt.Errorf("status %d: %s", resp.StatusCode, strings.TrimSpace(string(respBody))) - } - - var result SimilarAccountsResponse - if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { - return nil, 0, fmt.Errorf("decode response: %w", err) - } - return &result, time.Since(start), nil -} diff --git a/verify.py b/verify.py deleted file mode 100644 index 38224de..0000000 --- a/verify.py +++ /dev/null @@ -1,260 +0,0 @@ -#!/usr/bin/env python3 -"""Verify firehose embeddings match local EmbeddingGemma output. - -Reads JSONL events (same format as client -out), fetches each record via the -Bluesky API, prepares text identically to Divepool, embeds locally with -EmbeddingGemma, and compares via cosine similarity. Also shows three language -signals: firehose detected, locally confirmed (lingua), and author-declared. - -Usage: - HF_TOKEN=hf_... python3 verify.py events.jsonl # from client -out - -Requires HF_TOKEN env var (https://huggingface.co/google/embeddinggemma-300m). - -Precision strategy: Divepool runs EmbeddingGemma on M1 Mac Studios via MLX. -MLX stores weights in bf16 but computes all GEMM in float32 — Metal's -simdgroup_matrix uses the FP32 ALU pipeline and MLX kernels hardcode -AccumType=float. This script matches that: loads weights in bf16 for identical -rounding, then upcasts to float32 for compute. Expected cosine similarity -~0.995+. The remaining gap is from different attention implementations -(HuggingFace vs MLX Metal) and server-side batch padding context. -""" - -import json -import os -import sys -import tempfile -import urllib.request - -import numpy as np -import torch -from lingua import LanguageDetectorBuilder -from safetensors.torch import load_file -from transformers import AutoModel, AutoTokenizer - -PREFIXES = { - "cluster": "task: clustering | query: ", - "retrieval": "title: none | text: ", -} -MODEL_ID = "google/embeddinggemma-300m" - - -def prepare_text(record_value: dict, collection: str) -> str: - """Replicate Divepool's text preparation for posts and profiles.""" - if collection == "app.bsky.actor.profile": - return _prepare_profile_text(record_value) - return _prepare_post_text(record_value) - - -def _prepare_profile_text(record_value: dict) -> str: - """Replicate Divepool's ProfileText: displayName + description, joined by newline.""" - parts = [] - name = record_value.get("displayName", "") - if name: - # Replace newlines so \n reliably separates display name from description - name = name.replace("\n", " ").replace("\r", " ") - parts.append(name) - desc = record_value.get("description", "") - if desc: - parts.append(desc) - return "\n".join(parts).replace("\x00", "") - - -def _prepare_post_text(record_value: dict) -> str: - """Replicate Divepool's PostText: text + img/video alt texts + tags, joined by newline.""" - parts = [] - text = record_value.get("text", "") - if text: - parts.append(text) - - embed = record_value.get("embed") - if embed: - parts.extend(_extract_alt_texts(embed)) - - tags = record_value.get("tags") - if tags: - parts.append("tags: " + ", ".join(tags)) - - return "\n".join(parts).replace("\x00", "") - - -def _extract_alt_texts(embed: dict) -> list[str]: - alts = [] - if images := embed.get("images"): - for img in images: - if alt := img.get("alt", ""): - alts.append("img: " + alt) - if embed.get("$type", "").endswith("embed.video"): - if alt := embed.get("alt", ""): - alts.append("video: " + alt) - if media := embed.get("media"): - if images := media.get("images"): - for img in images: - if alt := img.get("alt", ""): - alts.append("img: " + alt) - if media.get("$type", "").endswith("embed.video"): - if alt := media.get("alt", ""): - alts.append("video: " + alt) - return alts - - -def fetch_record(did: str, collection: str, rkey: str) -> dict: - url = ( - f"https://public.api.bsky.app/xrpc/com.atproto.repo.getRecord" - f"?repo={did}&collection={collection}&rkey={rkey}" - ) - with urllib.request.urlopen(url) as resp: - return json.loads(resp.read()) - - -VALID_DIMS = {128, 768} - - -def validate_embedding(emb: np.ndarray, label: str) -> None: - """Reject embeddings that aren't 128d or 768d, or aren't L2-normalized.""" - if emb.shape[0] not in VALID_DIMS: - print(f" error: {label} has {emb.shape[0]}d, expected 128 or 768", file=sys.stderr, flush=True) - sys.exit(1) - norm = float(np.linalg.norm(emb)) - if abs(norm - 1.0) > 0.01: - print(f" error: {label} not L2-normalized (norm={norm:.4f})", file=sys.stderr, flush=True) - sys.exit(1) - - -def truncate_and_normalize(emb: np.ndarray, dim: int) -> np.ndarray: - """Matryoshka truncation: crop to dim, then L2-normalize.""" - cropped = emb[:dim] - return cropped / np.linalg.norm(cropped) - - -def cosine_similarity(a: np.ndarray, b: np.ndarray) -> float: - return float(np.dot(a, b) / (np.linalg.norm(a) * np.linalg.norm(b))) - - -def _download_dense_layer(token: str, path: str) -> str: - """Download a dense layer safetensors file from HuggingFace via HTTP. - - The huggingface_hub Python SDK sometimes hangs on xet-bridge CDN redirects, - so we use urllib as a reliable fallback. - """ - url = f"https://huggingface.co/{MODEL_ID}/resolve/main/{path}" - req = urllib.request.Request(url, headers={"Authorization": f"Bearer {token}"}) - tmp = os.path.join(tempfile.gettempdir(), f"embeddinggemma_{path.replace('/', '_')}") - if os.path.exists(tmp) and os.path.getsize(tmp) > 0: - return tmp - with urllib.request.urlopen(req) as resp, open(tmp, "wb") as f: - f.write(resp.read()) - return tmp - - -class EmbeddingModel: - """EmbeddingGemma: transformer + mean pooling + 2 dense layers + L2 normalize.""" - - def __init__(self, token: str): - self.tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, token=token) - # Load in bf16 for correct weight rounding, then upcast to float32 for compute. - # MLX on M1 does the same: bf16 storage, float32 GEMM (AccumType=float in Metal kernels). - self.model = AutoModel.from_pretrained( - MODEL_ID, token=token, trust_remote_code=True, dtype=torch.bfloat16, - ).float() - self.model.eval() - - d0_path = _download_dense_layer(token, "2_Dense/model.safetensors") - d1_path = _download_dense_layer(token, "3_Dense/model.safetensors") - # bf16→float32 round-trip matches server weight precision - self.dense0 = load_file(d0_path)["linear.weight"].to(torch.bfloat16).float() - self.dense1 = load_file(d1_path)["linear.weight"].to(torch.bfloat16).float() - - def encode(self, text: str, prefix: str) -> np.ndarray: - inputs = self.tokenizer( - prefix + text, - return_tensors="pt", - padding=True, - truncation=True, - max_length=2048, - ) - with torch.no_grad(): - out = self.model(**inputs) - - # Mean pooling (all float32 — model and weights already upcast in __init__) - h = out.last_hidden_state - mask = inputs["attention_mask"].unsqueeze(-1).float() - emb = (h * mask).sum(dim=1) / mask.sum(dim=1) - emb = emb.squeeze() - - # Dense projection (identity activation, no bias) - emb = emb @ self.dense0.T - emb = emb @ self.dense1.T - - # L2 normalize - emb = emb / emb.norm() - return emb.numpy().astype(np.float32) - - -def main(): - token = os.environ.get("HF_TOKEN") - if not token: - print("error: HF_TOKEN env var required (https://huggingface.co/settings/tokens)", file=sys.stderr) - sys.exit(1) - - if len(sys.argv) < 2: - print("usage: verify.py ", file=sys.stderr) - sys.exit(1) - - with open(sys.argv[1]) as f: - items = [json.loads(line) for line in f if line.strip()] - - print("Loading lingua detector...", flush=True) - lang_detector = LanguageDetectorBuilder.from_all_languages().build() - - print(f"Loading {MODEL_ID}...", flush=True) - emb_model = EmbeddingModel(token) - print(f"Loaded. Verifying {len(items)} item(s).\n", flush=True) - - for item in items: - did, col, rkey = item["did"], item["col"], item["rkey"] - uri = f"at://{did}/{col}/{rkey}" - print(f"{uri}", flush=True) - - try: - record = fetch_record(did, col, rkey) - except Exception as e: - print(f" error: fetch failed ({e})\n", flush=True) - continue - - value = record.get("value", {}) - text = prepare_text(value, col) - if not text: - print(f" skip: no text\n", flush=True) - continue - - print(f" text: {text[:80]}", flush=True) - - # Language signals: firehose detected, lingua confirmed, author declared - firehose_lang = item.get("lang", "") - lingua_result = lang_detector.detect_language_of(text) - lingua_lang = lingua_result.iso_code_639_1.name.lower() if lingua_result else "?" - declared_langs = value.get("langs", []) - declared = ",".join(declared_langs) if declared_langs else "-" - print(f" lang: firehose={firehose_lang} lingua={lingua_lang} declared={declared}", flush=True) - - for key, task in [("c", "cluster"), ("r", "retrieval")]: - expected = item.get(key) - if not expected: - continue - expected = np.array(expected, dtype=np.float32) - validate_embedding(expected, f"{task} embedding") - dim = expected.shape[0] - - local = emb_model.encode(text, PREFIXES[task]) - if dim < local.shape[0]: - local = truncate_and_normalize(local, dim) - - sim = cosine_similarity(expected, local) - print(f" {task:10s} {dim}d cosine similarity: {sim:.6f}", flush=True) - - print(flush=True) - - -if __name__ == "__main__": - main()