File size: 8,264 Bytes
7e9cfd1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
"""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()