stanceeval2026 / code /src /gen_synth.py
zaher-m's picture
Add files using upload-large-folder tool
7e9cfd1 verified
Raw
History Blame Contribute Delete
8.26 kB
"""Have a chat LLM make up labeled tweets for a target -- handy for a dev set
on unseen targets, or to pad out the retrieval pool. Varies dialect and tone
for a bit of diversity.
python -m src.gen_synth --target "Women Driving" \\
--desc "قيادة المرأة للسيارة" --per_class 150 \\
--out data/synth/women_driving.csv
"""
import argparse
import json
import os
import re
import urllib.request
from concurrent.futures import ThreadPoolExecutor
import numpy as np
import pandas as pd
STANCE_AR = {
"Favor": "مؤيدة وداعمة",
"Against": "معارضة ورافضة",
"None": "لا تُظهر موقفاً واضحاً (خبر أو سؤال أو تعليق محايد)",
}
DIALECTS = ["سعودية خليجية", "مصرية", "شامية", "عربية فصحى"]
TONES = ["جادة", "ساخرة تهكمية", "على شكل سؤال", "انفعالية عاطفية",
"عامية عفوية"]
SYSTEM = (
"أنت مستخدم عربي على منصة إكس (تويتر). اكتب تغريدة واحدة واقعية قصيرة "
"كما يكتبها الناس فعلاً: عفوية، قد تحوي وسماً أو إيموجي أو اختصاراً، "
"بلا ترقيم زائد. اكتب التغريدة فقط دون أي شرح أو علامات اقتباس."
)
SYSTEM_HARD = (
"أنت مستخدم عربي على منصة إكس (تويتر). اكتب تغريدة واحدة واقعية يصعب "
"تصنيف موقفها بسهولة: عبّر عن الموقف بشكل غير مباشر أو ضمني، استخدم "
"السخرية والتلميح والمبالغة واللهجة العامية والأخطاء الإملائية، وقد لا "
"تذكر الهدف صراحة. اكتب التغريدة فقط دون شرح أو علامات اقتباس."
)
SYSTEM_STYLE = (
"أنت مستخدم سعودي على منصة إكس (تويتر). سأعطيك أمثلة حقيقية لتغريدات "
"حول موضوع معيّن. اكتب تغريدة جديدة تحاكي أسلوبها ولهجتها ووسومها "
"وطريقة كتابتها (السخرية، الإيموجي، تطويل الحروف، الأخطاء الإملائية) "
"لكنها تعبّر بوضوح عن الموقف المطلوب. اكتب التغريدة فقط دون شرح أو "
"علامات اقتباس."
)
def build_msgs_style(target, desc, stance, anchors, tag=None):
tgt = f"{target} ({desc})" if desc else target
ex = "\n".join(f"- {a}" for a in anchors)
tag_line = f"استعمل الوسم {tag} إن ناسب. " if tag else ""
user = (
f"أمثلة على أسلوب التغريدات حول {tgt}:\n{ex}\n\n"
f"اكتب تغريدة جديدة موقفها {STANCE_AR[stance]} تجاه {tgt} بنفس "
f"الأسلوب واللهجة والطابع. {tag_line}اجعلها واقعية وغير مكررة."
)
return [{"role": "system", "content": SYSTEM_STYLE},
{"role": "user", "content": user}]
def build_msgs(target, desc, stance, dialect, tone, hard=False):
tgt = f"{target} ({desc})" if desc else target
if hard:
user = (
f"اكتب تغريدة موقفها {STANCE_AR[stance]} تجاه موضوع: {tgt}، "
f"لكن بشكل غير مباشر وضمني وصعب. اللهجة: {dialect}. "
f"الأسلوب: {tone}. اجعلها واقعية وغير مكررة."
)
return [{"role": "system", "content": SYSTEM_HARD},
{"role": "user", "content": user}]
user = (
f"اكتب تغريدة {STANCE_AR[stance]} تجاه: {tgt}.\n"
f"اللهجة: {dialect}. الأسلوب: {tone}.\n"
"اجعلها مختلفة وغير مكررة وواقعية."
)
return [{"role": "system", "content": SYSTEM},
{"role": "user", "content": user}]
def call(base_url, model, msgs, temperature, max_tokens=90, timeout=180):
body = json.dumps({
"model": model, "temperature": temperature, "top_p": 0.95,
"max_tokens": max_tokens, "messages": msgs,
}).encode("utf-8")
req = urllib.request.Request(
base_url.rstrip("/") + "/chat/completions", data=body,
headers={"Content-Type": "application/json"},
)
with urllib.request.urlopen(req, timeout=timeout) as r:
return json.load(r)["choices"][0]["message"]["content"]
def clean(text):
t = text.strip().strip('"').strip("«»").strip()
t = re.sub(r"\s+", " ", t)
return t
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--target", required=True)
ap.add_argument("--desc", default="")
ap.add_argument("--per_class", type=int, default=150)
ap.add_argument("--out", required=True)
ap.add_argument("--hard", action="store_true",
help="generate subtle/implicit hard-to-classify tweets")
ap.add_argument("--style_file", default=None,
help="CSV of real (unlabeled) tweets used as style "
"anchors; switches to style-grounded generation")
ap.add_argument("--style_col", default="tweet_text")
ap.add_argument("--seed_tags", default="",
help="comma-separated hashtags to seed into prompts")
ap.add_argument("--n_anchors", type=int, default=3)
ap.add_argument("--temperature", type=float, default=1.0)
ap.add_argument("--concurrency", type=int, default=24)
ap.add_argument("--base_url", default=os.environ.get("AUG_BASE_URL", ""))
ap.add_argument("--model", default=os.environ.get("AUG_MODEL", ""))
args = ap.parse_args()
if not args.base_url or not args.model:
raise SystemExit("set AUG_BASE_URL/AUG_MODEL")
anchors_pool, tags = [], []
if args.style_file:
sdf = pd.read_csv(args.style_file)
anchors_pool = [str(t) for t in sdf[args.style_col].tolist()
if isinstance(t, str) or not pd.isna(t)]
if args.seed_tags:
tags = [t.strip() for t in args.seed_tags.split(",") if t.strip()]
rng = np.random.RandomState(0)
jobs = []
for stance in ["Favor", "Against", "None"]:
for i in range(args.per_class):
if anchors_pool:
idx = rng.choice(len(anchors_pool),
size=min(args.n_anchors, len(anchors_pool)),
replace=False)
anchors = [anchors_pool[k] for k in idx]
tag = tags[i % len(tags)] if tags else None
jobs.append((stance, "style", anchors, tag))
else:
dialect = DIALECTS[i % len(DIALECTS)]
tone = TONES[i % len(TONES)]
jobs.append((stance, "plain", dialect, tone))
def work(job):
stance = job[0]
try:
if job[1] == "style":
_, _, anchors, tag = job
msgs = build_msgs_style(args.target, args.desc, stance,
anchors, tag)
else:
_, _, dialect, tone = job
msgs = build_msgs(args.target, args.desc, stance,
dialect, tone, args.hard)
out = call(args.base_url, args.model, msgs, args.temperature)
return stance, clean(out)
except Exception:
return stance, ""
rows = []
with ThreadPoolExecutor(max_workers=args.concurrency) as ex:
for j, (stance, text) in enumerate(ex.map(work, jobs)):
if text and len(text) > 8:
rows.append({"text": text, "target": args.target,
"stance": stance})
if (j + 1) % 100 == 0:
print(f" {j + 1}/{len(jobs)}", flush=True)
df = pd.DataFrame(rows).drop_duplicates(subset=["text"])
df = df.sample(frac=1.0, random_state=np.random.RandomState(0))
os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True)
df.to_csv(args.out, index=False, encoding="utf-8-sig")
print(f"[write] {len(df)} rows -> {args.out}")
print(df["stance"].value_counts().to_dict())
if __name__ == "__main__":
main()