-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathbaseline.py
More file actions
147 lines (118 loc) · 4.8 KB
/
Copy pathbaseline.py
File metadata and controls
147 lines (118 loc) · 4.8 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
"""
SemEval 2026 Task 4 Subtask 1 Baseline
Narrative Similarity: Given an anchor text, predict which of two candidate texts
(A or B) is more similar to the anchor.
Uses Gemini 3 Flash Preview with reasoning in schema for predictions.
"""
import os
import pandas as pd
from google import genai
from google.genai import types
import json
import asyncio
import time
from tqdm.asyncio import tqdm
# Configuration
API_KEY = os.environ.get("GEMINI_API_KEY")
if not API_KEY:
raise SystemExit("GEMINI_API_KEY is required. Copy .env.example to .env, fill a rotated key, and load it into your shell.")
MODEL = "gemini-3-flash-preview"
DATA_PATH = "data/dev_track_a.jsonl"
OUTPUT_PATH = "data/baseline_results.csv"
CONCURRENCY_LIMIT = 50
client = genai.Client(api_key=API_KEY)
# Schema for structured JSON output with reasoning
response_schema = types.Schema(
type=types.Type.OBJECT,
properties={
"reasoning": types.Schema(
type=types.Type.STRING,
description="Your step-by-step analysis comparing the stories based on abstract themes, course of action, and outcomes. Do NOT focus on surface features like setting, era, or genre."
),
"decision": types.Schema(
type=types.Type.STRING,
enum=["A", "B"],
description="Your final choice: A if Story A is more similar to Anchor, B if Story B is more similar."
)
},
required=["reasoning", "decision"]
)
async def predict(row: dict, semaphore: asyncio.Semaphore) -> tuple[str, str]:
"""Make a single prediction for which text is more similar to the anchor.
Returns:
tuple: (decision, reasoning)
"""
prompt = f"""Which story (A or B) is more similar to the Anchor story?
Focus on NARRATIVE SIMILARITY based on:
1. Abstract themes (core problems, central ideas) - IGNORE concrete settings
2. Course of action (sequence of events)
3. Outcomes (how the story ends)
Anchor: {row['anchor_text']}
Story A: {row['text_a']}
Story B: {row['text_b']}
Provide your reasoning and decision."""
async with semaphore:
for attempt in range(3):
try:
response = await client.aio.models.generate_content(
model=MODEL,
contents=prompt,
config=types.GenerateContentConfig(
response_mime_type="application/json",
response_schema=response_schema
)
)
data = json.loads(response.text)
return data.get("decision", "A"), data.get("reasoning", "")
except Exception as e:
if attempt < 2:
await asyncio.sleep(2 ** attempt)
return "A", "" # Default fallback
async def main():
print(f"{'═'*60}")
print(f"SemEval 2026 Task 4 Subtask 1 - Baseline")
print(f"Model: {MODEL} (with reasoning in schema)")
print(f"{'═'*60}")
# Load dataset
print("\n📂 Loading dataset...")
if not os.path.exists(DATA_PATH):
raise SystemExit(f"data file not found: {DATA_PATH}. See data/README.md for download and placement notes.")
df = pd.read_json(DATA_PATH, lines=True).head(200)
print(f" Evaluating on {len(df)} examples from {DATA_PATH}")
# Run predictions
print(f"\n🧠 Running predictions with reasoning...")
start_time = time.time()
semaphore = asyncio.Semaphore(CONCURRENCY_LIMIT)
tasks = [predict(row, semaphore) for _, row in df.iterrows()]
results = await tqdm.gather(*tasks)
elapsed = time.time() - start_time
# Unpack predictions and reasoning
predictions = [r[0] for r in results]
reasoning = [r[1] for r in results]
# Store results
df["predicted_decision"] = predictions
df["predicted_text_a_is_closer"] = [p == "A" for p in predictions]
df["reasoning"] = reasoning
# Calculate accuracy
correct = (df["predicted_text_a_is_closer"] == df["text_a_is_closer"]).sum()
accuracy = correct / len(df)
# Results
print(f"\n{'═'*60}")
print(f"📊 RESULTS")
print(f"{'─'*60}")
print(f"Accuracy: {accuracy:.1%} ({correct}/{len(df)})")
print(f"Time: {elapsed:.1f}s ({elapsed/len(df):.2f}s per example)")
print(f"{'═'*60}")
# Show a sample reasoning
print(f"\n💭 Sample reasoning (first example):")
for i, r in enumerate(reasoning):
if r:
print(f"{'─'*60}")
print(f"Example {i+1}: Predicted {predictions[i]}, Actual {'A' if df.iloc[i]['text_a_is_closer'] else 'B'}")
print(f"Reasoning: {r[:800]}..." if len(r) > 800 else f"Reasoning: {r}")
break
# Save results
df.to_csv(OUTPUT_PATH, index=False)
print(f"\n📂 Results saved to: {OUTPUT_PATH}")
if __name__ == "__main__":
asyncio.run(main())