-
Notifications
You must be signed in to change notification settings - Fork 33
Add Market: poll a simulated audience on which option it prefers #136
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,163 @@ | ||
| #!/usr/bin/env python3 | ||
| """Load the Market persona corpus into Postgres. | ||
|
|
||
| Reads Nemotron-Personas-USA parquet shards (CC-BY-4.0, NVIDIA), renders each | ||
| row into a ~420-character panel text plus structured attributes, embeds the | ||
| text with OpenAI text-embedding-3-small at 512 dimensions through OpenRouter, | ||
| and bulk-inserts into market_personas. Resumable: rows whose id is already | ||
| present are skipped, so a crashed run restarts where it stopped. | ||
|
|
||
| Env: MARKET_DATABASE_URL_DIRECT (or DATABASE_URL), OPENROUTER_EMBED_KEY. | ||
| Usage: python3 market-corpus.py <shard.parquet> [<shard.parquet> ...] | ||
| """ | ||
| import json, os, sys, time, urllib.request, urllib.error | ||
| from concurrent.futures import ThreadPoolExecutor | ||
|
|
||
| import duckdb | ||
| import psycopg | ||
|
|
||
| DB = os.environ.get("MARKET_DATABASE_URL_DIRECT") or os.environ["DATABASE_URL"] | ||
| KEY = os.environ["OPENROUTER_EMBED_KEY"] | ||
| EMBED_URL = "https://openrouter.ai/api/v1/embeddings" | ||
| MODEL = "openai/text-embedding-3-small" | ||
| DIMS = 512 | ||
| BATCH = 384 # texts per embedding request | ||
| WORKERS = 6 # concurrent embedding requests | ||
| ROWS_PER_SHARD = 1_000_000 # id space per shard; ids are shard_index * this + row | ||
|
|
||
| REGION = { # census regions, for the segments table | ||
| "Connecticut":"Northeast","Maine":"Northeast","Massachusetts":"Northeast","New Hampshire":"Northeast", | ||
| "Rhode Island":"Northeast","Vermont":"Northeast","New Jersey":"Northeast","New York":"Northeast","Pennsylvania":"Northeast", | ||
| "Illinois":"Midwest","Indiana":"Midwest","Michigan":"Midwest","Ohio":"Midwest","Wisconsin":"Midwest", | ||
| "Iowa":"Midwest","Kansas":"Midwest","Minnesota":"Midwest","Missouri":"Midwest","Nebraska":"Midwest", | ||
| "North Dakota":"Midwest","South Dakota":"Midwest", | ||
| "Delaware":"South","Florida":"South","Georgia":"South","Maryland":"South","North Carolina":"South", | ||
| "South Carolina":"South","Virginia":"South","District of Columbia":"South","West Virginia":"South", | ||
| "Alabama":"South","Kentucky":"South","Mississippi":"South","Tennessee":"South", | ||
| "Arkansas":"South","Louisiana":"South","Oklahoma":"South","Texas":"South", | ||
| "Arizona":"West","Colorado":"West","Idaho":"West","Montana":"West","Nevada":"West", | ||
| "New Mexico":"West","Utah":"West","Wyoming":"West","Alaska":"West","California":"West", | ||
| "Hawaii":"West","Oregon":"West","Washington":"West", | ||
| } | ||
|
|
||
| def age_band(age): | ||
| for lo, hi in ((18,24),(25,34),(35,44),(45,54),(55,64),(65,120)): | ||
| if lo <= age <= hi: | ||
| return f"{lo}-{hi}" if hi < 120 else "65+" | ||
| return "under-18" | ||
|
|
||
| def clip_sentence(text, limit): | ||
| """Trim at the last sentence boundary within limit; hard-cut as a fallback.""" | ||
| if not text or len(text) <= limit: | ||
| return text or "" | ||
| cut = text[:limit] | ||
| dot = cut.rfind(". ") | ||
| return cut[: dot + 1] if dot > limit // 2 else cut | ||
|
|
||
| MARITAL = {"married_present": "married", "married_absent": "married, spouse away", | ||
| "never_married": "single", "divorced": "divorced", "widowed": "widowed", "separated": "separated"} | ||
|
|
||
| def humanize(value): | ||
| return (value or "").replace("_", " ").strip() | ||
|
|
||
| def render(row): | ||
| (persona, prof, sex, age, marital, edu, field, occ, city, state, hobbies) = row | ||
| bits = [f"{age}-year-old {humanize(sex).lower()}, {MARITAL.get(marital, humanize(marital))}."] | ||
| edu_txt = humanize(edu) or "unknown education" | ||
| if field: | ||
| edu_txt += f" in {humanize(field)}" | ||
| bits.append(f"Occupation: {humanize(occ) or 'unknown'} ({edu_txt}).") | ||
| bits.append(f"Lives in {city}, {state}.") | ||
| bits.append(clip_sentence(prof or persona or "", 240)) | ||
| if hobbies: | ||
| take = [h.strip() for h in hobbies[:3] if h and h.strip()] | ||
| if take: | ||
| bits.append("Interests: " + "; ".join(take) + ".") | ||
|
Comment on lines
+73
to
+75
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🗄️ Data Integrity & Integration | 🟠 Major | ⚡ Quick win Parse the hobbies field before selecting interests. The source Parquet schema stores 🤖 Prompt for AI Agents |
||
| text = " ".join(b for b in bits if b) | ||
| attrs = { | ||
| "age": age, "age_band": age_band(age), "sex": humanize(sex).lower(), | ||
| "marital": MARITAL.get(marital, humanize(marital)), | ||
| "education": humanize(edu), "occupation": humanize(occ), "state": state, | ||
| "region": REGION.get(state, "Other"), | ||
| } | ||
| return text, attrs | ||
|
|
||
| def embed(texts, tries=6): | ||
| body = json.dumps({"model": MODEL, "input": texts, "dimensions": DIMS}).encode() | ||
| for attempt in range(tries): | ||
| req = urllib.request.Request(EMBED_URL, data=body, headers={ | ||
| "Authorization": f"Bearer {KEY}", "Content-Type": "application/json"}) | ||
| try: | ||
| with urllib.request.urlopen(req, timeout=120) as r: | ||
| data = json.load(r) | ||
| vecs = [d["embedding"] for d in sorted(data["data"], key=lambda d: d["index"])] | ||
| if len(vecs) != len(texts) or any(len(v) != DIMS for v in vecs): | ||
| raise ValueError("embedding response shape mismatch") | ||
| return vecs | ||
| except (urllib.error.HTTPError, urllib.error.URLError, ValueError, TimeoutError, OSError) as e: | ||
| status = getattr(e, "code", None) | ||
| if attempt == tries - 1: | ||
| raise | ||
| time.sleep(min(2 ** attempt + 1, 30) if status in (429, None) else 2) | ||
|
|
||
| def main(shards): | ||
| with psycopg.connect(DB) as check: | ||
| done = {r[0] for r in check.execute("SELECT id FROM market_personas").fetchall()} | ||
| print(f"resume: {len(done)} rows already loaded", flush=True) | ||
|
|
||
| conn = psycopg.connect(DB, autocommit=True) | ||
| inserted = 0 | ||
| started = time.time() | ||
| pool = ThreadPoolExecutor(max_workers=WORKERS) | ||
| for si, shard in enumerate(shards): | ||
| base = si * ROWS_PER_SHARD | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🗄️ Data Integrity & Integration | 🟠 Major | 🏗️ Heavy lift Keep persona IDs stable across resumed runs.
🤖 Prompt for AI Agents |
||
| rel = duckdb.connect().execute(f""" | ||
| SELECT persona, professional_persona, sex, age, marital_status, | ||
| education_level, bachelors_field, occupation, city, state, | ||
| hobbies_and_interests_list | ||
| FROM read_parquet('{shard}') | ||
| WHERE age >= 18 | ||
| {"LIMIT " + os.environ["MARKET_MAX_ROWS"] if os.environ.get("MARKET_MAX_ROWS") else ""}""") | ||
| batch_rows, batch_ids, futures = [], [], [] | ||
|
|
||
| def flush(rows, ids): | ||
| texts_attrs = [render(r) for r in rows] | ||
| texts = [t for t, _ in texts_attrs] | ||
| vecs = embed(texts) | ||
| payload = [ | ||
| (pid, text, json.dumps(attrs), "[" + ",".join(f"{x:.5f}" for x in vec) + "]") | ||
| for pid, (text, attrs), vec in zip(ids, texts_attrs, vecs) | ||
| ] | ||
| with conn.cursor() as cur: | ||
| cur.executemany( | ||
| "INSERT INTO market_personas (id, panel_text, attrs, embedding) " | ||
| "VALUES (%s, %s, %s, %s) ON CONFLICT (id) DO NOTHING", payload) | ||
| return len(payload) | ||
|
|
||
| row_index = 0 | ||
| while True: | ||
| chunk = rel.fetchmany(BATCH) | ||
| if not chunk: | ||
| break | ||
| ids = list(range(base + row_index, base + row_index + len(chunk))) | ||
| row_index += len(chunk) | ||
| keep = [(r, i) for r, i in zip(chunk, ids) if i not in done] | ||
| if not keep: | ||
| continue | ||
| rows = [r for r, _ in keep] | ||
| kept_ids = [i for _, i in keep] | ||
| futures.append(pool.submit(flush, rows, kept_ids)) | ||
| if len(futures) >= WORKERS * 2: | ||
| for f in futures: | ||
| inserted += f.result() | ||
| futures = [] | ||
| rate = inserted / max(time.time() - started, 1) | ||
| print(f"shard {si}: {row_index} read, {inserted} inserted total, {rate:.0f} rows/s", flush=True) | ||
| for f in futures: | ||
| inserted += f.result() | ||
| print(f"shard {si} complete: {row_index} rows read", flush=True) | ||
| pool.shutdown() | ||
| print(f"DONE: {inserted} inserted in {time.time()-started:.0f}s", flush=True) | ||
|
|
||
| if __name__ == "__main__": | ||
| main(sys.argv[1:]) | ||
| Original file line number | Diff line number | Diff line change | ||||||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| @@ -0,0 +1,108 @@ | ||||||||||||||||
| /** | ||||||||||||||||
| * Live Market pipeline harness: retrieval → membership → panel → votes, | ||||||||||||||||
| * against the real corpus and real Jev, without the worker or billing. | ||||||||||||||||
| * | ||||||||||||||||
| * bun --env-file=.dev.vars scripts/market-live.ts "<audience>" "<option a>" "<option b>" [population] [decision] | ||||||||||||||||
| * | ||||||||||||||||
| * Requires DATABASE_URL (market corpus), TYPESAFE_API_KEY, and — for vector | ||||||||||||||||
| * retrieval — OPENROUTER_EMBED_KEY in the environment. | ||||||||||||||||
| */ | ||||||||||||||||
| import { neon } from "@neondatabase/serverless"; | ||||||||||||||||
| import { jevKeys } from "../src/jev"; | ||||||||||||||||
| import { newMeter } from "../src/cost"; | ||||||||||||||||
| import { | ||||||||||||||||
| aggregate, | ||||||||||||||||
| readMarketRequest, | ||||||||||||||||
| runVotes, | ||||||||||||||||
| samplePanel, | ||||||||||||||||
| scoreMembership, | ||||||||||||||||
| seededRandom, | ||||||||||||||||
| sha256Hex, | ||||||||||||||||
| voteBatches, | ||||||||||||||||
| MARKET_CORPUS_VERSION, | ||||||||||||||||
| MEMBERSHIP_SHORTLIST, | ||||||||||||||||
| type PanelMember, | ||||||||||||||||
| } from "../src/market"; | ||||||||||||||||
|
|
||||||||||||||||
| const [audience, a, b, populationRaw, decisionRaw] = process.argv.slice(2); | ||||||||||||||||
| if (!audience || !a || !b) { | ||||||||||||||||
| console.error('usage: bun scripts/market-live.ts "<audience>" "<option a>" "<option b>" [population] [decision]'); | ||||||||||||||||
| process.exit(1); | ||||||||||||||||
| } | ||||||||||||||||
| const request = readMarketRequest({ | ||||||||||||||||
| audience, | ||||||||||||||||
| options: [a, b], | ||||||||||||||||
| population: populationRaw ? Number(populationRaw) : 200, | ||||||||||||||||
| ...(decisionRaw ? { decision: decisionRaw } : {}), | ||||||||||||||||
| }); | ||||||||||||||||
|
|
||||||||||||||||
| const sql = neon(process.env.MARKET_DATABASE_URL ?? process.env.DATABASE_URL!); | ||||||||||||||||
| const keys = jevKeys(process.env as Record<string, string>)!; | ||||||||||||||||
|
Comment on lines
+39
to
+40
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🩺 Stability & Availability | 🟡 Minor | ⚡ Quick win Validate the required environment before you connect. If Proposed fix-const sql = neon(process.env.MARKET_DATABASE_URL ?? process.env.DATABASE_URL!);
-const keys = jevKeys(process.env as Record<string, string>)!;
+const url = process.env.MARKET_DATABASE_URL ?? process.env.DATABASE_URL;
+if (!url) { console.error("Set MARKET_DATABASE_URL or DATABASE_URL."); process.exit(1); }
+const sql = neon(url);
+const keys = jevKeys(process.env as Record<string, string>);
+if (!keys) { console.error("Set TYPESAFE_API_KEY (or another Jev provider key)."); process.exit(1); }Based on learnings: "explicitly validate they are defined ... rather than using non-null assertions". 📝 Committable suggestion
Suggested change
🤖 Prompt for AI AgentsSource: Learnings |
||||||||||||||||
| const meter = newMeter(); | ||||||||||||||||
| const t0 = Date.now(); | ||||||||||||||||
| const lap = (label: string) => console.log(`${label}: ${((Date.now() - t0) / 1000).toFixed(1)}s, spend $${meter.usd.toFixed(4)}`); | ||||||||||||||||
|
|
||||||||||||||||
| // 1. Embed the audience (optional). | ||||||||||||||||
| let embedding: string | null = null; | ||||||||||||||||
| if (process.env.OPENROUTER_EMBED_KEY) { | ||||||||||||||||
| const response = await fetch("https://openrouter.ai/api/v1/embeddings", { | ||||||||||||||||
| method: "POST", | ||||||||||||||||
| headers: { authorization: `Bearer ${process.env.OPENROUTER_EMBED_KEY}`, "content-type": "application/json" }, | ||||||||||||||||
| body: JSON.stringify({ model: "openai/text-embedding-3-small", input: [request.audience], dimensions: 512 }), | ||||||||||||||||
| }); | ||||||||||||||||
| const data = await response.json() as { data?: { embedding: number[] }[] }; | ||||||||||||||||
| if (data.data?.[0]) embedding = `[${data.data[0].embedding.map((x) => x.toFixed(5)).join(",")}]`; | ||||||||||||||||
| } | ||||||||||||||||
| lap(`embed (${embedding ? "ok" : "SKIPPED"})`); | ||||||||||||||||
|
|
||||||||||||||||
| // 2. Retrieve. | ||||||||||||||||
| const shortlist = (embedding | ||||||||||||||||
| ? (await sql.transaction((tx) => [ | ||||||||||||||||
| tx`SET LOCAL hnsw.ef_search = 1000`, | ||||||||||||||||
| tx` | ||||||||||||||||
| WITH vec AS (SELECT id FROM market_personas ORDER BY embedding <=> ${embedding}::halfvec(512) LIMIT 1000), | ||||||||||||||||
| kw AS (SELECT id FROM market_personas | ||||||||||||||||
| WHERE tsv @@ websearch_to_tsquery('english', ${request.audience}) | ||||||||||||||||
| ORDER BY ts_rank(tsv, websearch_to_tsquery('english', ${request.audience})) DESC LIMIT 1400) | ||||||||||||||||
| SELECT p.id, p.panel_text FROM market_personas p | ||||||||||||||||
| WHERE p.id IN (SELECT id FROM vec UNION SELECT id FROM kw)`, | ||||||||||||||||
| ]))[1] | ||||||||||||||||
| : await sql` | ||||||||||||||||
| SELECT id, panel_text FROM market_personas | ||||||||||||||||
| WHERE tsv @@ websearch_to_tsquery('english', ${request.audience}) | ||||||||||||||||
| ORDER BY ts_rank(tsv, websearch_to_tsquery('english', ${request.audience})) DESC LIMIT 2200` | ||||||||||||||||
| ) as { id: number; panel_text: string }[]; | ||||||||||||||||
| lap(`retrieve (${shortlist.length} candidates)`); | ||||||||||||||||
|
|
||||||||||||||||
| // 3. Membership scoring. | ||||||||||||||||
| const audienceId = await sha256Hex(`${MARKET_CORPUS_VERSION}\n${request.audience}\n${request.population}`); | ||||||||||||||||
| const random = seededRandom(audienceId); | ||||||||||||||||
| const scored = shortlist | ||||||||||||||||
| .map((row) => ({ row, key: random() })) | ||||||||||||||||
| .sort((x, y) => x.key - y.key) | ||||||||||||||||
| .slice(0, MEMBERSHIP_SHORTLIST) | ||||||||||||||||
| .map(({ row }) => ({ id: row.id, text: row.panel_text })); | ||||||||||||||||
| const weights = await scoreMembership(keys, request.audience, scored, meter); | ||||||||||||||||
| const histogram = [0, 0, 0, 0, 0]; | ||||||||||||||||
| for (const w of weights) histogram[Math.min(4, Math.floor(w * 4))]++; | ||||||||||||||||
| lap(`membership (weights 0-.25/.25-.5/.5-.75/.75-1/1: ${histogram.join("/")})`); | ||||||||||||||||
|
|
||||||||||||||||
| // 4. Panel. | ||||||||||||||||
| const members = scored.map((candidate, i) => ({ id: candidate.id, weight: weights[i] })).filter((m) => m.weight > 0); | ||||||||||||||||
| const panelIds = samplePanel(members, request.population, audienceId); | ||||||||||||||||
| const textById = new Map(scored.map((candidate) => [candidate.id, candidate.text])); | ||||||||||||||||
| const attrsRows = await sql`SELECT id, attrs FROM market_personas WHERE id = ANY(${panelIds.map((m) => m.id)})` as { id: number; attrs: Record<string, unknown> }[]; | ||||||||||||||||
| const attrsById = new Map(attrsRows.map((row) => [row.id, row.attrs])); | ||||||||||||||||
| const panel = new Map<number, PanelMember>(panelIds.map((m) => [m.id, { id: m.id, weight: m.weight, attrs: attrsById.get(m.id) ?? {} }])); | ||||||||||||||||
| console.log(`panel: ${panelIds.length} of ${members.length} eligible`); | ||||||||||||||||
| for (const m of panelIds.slice(0, 3)) console.log(` · w=${m.weight.toFixed(2)} ${textById.get(m.id)?.slice(0, 130)}`); | ||||||||||||||||
|
|
||||||||||||||||
| // 5. Votes. | ||||||||||||||||
| const batches = voteBatches(request.options, panelIds.map((m) => ({ id: m.id, text: textById.get(m.id)! })), request.decision); | ||||||||||||||||
| const { votes, failedRespondents } = await runVotes(keys, batches, request.options, meter); | ||||||||||||||||
| lap(`votes (${votes.length} answered, ${failedRespondents} failed)`); | ||||||||||||||||
|
|
||||||||||||||||
| // 6. Aggregate. | ||||||||||||||||
| const result = aggregate(request, audienceId, scored.length, panel, votes); | ||||||||||||||||
| console.log(JSON.stringify(result, null, 1)); | ||||||||||||||||
| lap("total"); | ||||||||||||||||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
🎯 Functional Correctness | 🟠 Major | ⚡ Quick win
Remove the deployed market secret when the option is disabled.
If
MARKET_DATABASE_URLwas deployed previously, omitting it from this file does not remove it from the Worker. Wrangler preserves existing secrets that are absent from--secrets-file. Removing the GitHub secret therefore leaves the endpoint connected to the old database instead of making it return 503. Explicitly remove the deployed secret when this option is disabled. (developers.cloudflare.com)🤖 Prompt for AI Agents