File size: 6,636 Bytes
30ae36c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import os
import glob
import pandas as pd
import argparse
from google import genai
from tqdm import tqdm
import time
import re
from word_segmentation_vi import word_segmentation_vi

def setup_genai(api_key):
    """Configure the Google Generative AI client with your API key"""
    return genai.Client(api_key=api_key)

def classify_text(model, text):
    """Classify Vietnamese text into hate speech categories using Google's Generative AI"""
    prompt = f"""
    Analyze the following Vietnamese text for hate speech (each sentence is separated by a newline):
    "{text}"
    
    Rate it on these categories (0=NORMAL, 1=CLEAN, 2=OFFENSIVE, 3=HATE):
    - individual (targeting specific individuals)
    - groups (targeting groups or organizations)
    - religion/creed (targeting religious groups or beliefs)
    - race/ethnicity (racial/ethnic hate speech)
    - politics (political hate speech)
    If the text doesn't specify a person or group in a category, return 0 for that category.
    Else, return 1 for CLEAN, 2 for OFFENSIVE, or 3 for HATE.
    
    For each sentence in the text, return only 5 numbers separated by commas (corresponding to the label of individual, groups, religion/creed, race/ethnicity, politics) and numbers for each sentence seperated by newlines, like (with no other text): 
    0,1,0,0,0
    1,0,0,0,2
    """
    
    try:
        response = model.models.generate_content(model="gemini-2.0-flash", contents=prompt)
        values = response.text.strip().split('\n')
        values = [line.split(',') for line in values]
        return values
        
    except Exception as e:
        print(f"Error classifying text: {e}")
        return None

def process_file(input_file, output_file, model, rate_limit_pause=4):
    """Process a single CSV file to match the test.csv format"""
    print(f"Processing {input_file}...")
    
    # Read the input file
    try:
        df = pd.read_csv(input_file)
    except Exception as e:
        print(f"Error reading {input_file}: {e}")
        return
    
    # Rename column free_text to content
    if 'free_text' in df.columns:
        df.rename(columns={'free_text': 'content'}, inplace=True)
    elif 'content' not in df.columns:
        print(f"Error: 'content' column not found in {input_file}")
        return
    
    # Ensure all required columns exist
    category_columns = ['individual', 'groups', 'religion/creed', 'race/ethnicity', 'politics']
    for col in category_columns:
        if col not in df.columns:
            # Change column type to int if it doesn't exist
            df[col] = 0
    
    # Process each batch (100 rows at a time)
    batch_size = 100
    for start in tqdm(range(0, len(df), batch_size), desc="Processing batches"):
        end = min(start + batch_size, len(df))
        batch_df = df.iloc[start:end]
        
        # Skip if all categories already have values
        if all(batch_df[cat].all() != 0 for cat in category_columns):
            continue
        
        # Join 50 rows by newlines, and classify all at once
        text_to_classify = "\n".join([str(sentence) for sentence in batch_df['content'].tolist()])
        classifications = classify_text(model, text_to_classify)

        # Try 2 more times, else skip
        if classifications is None:
            for _ in range(2):
                classifications = classify_text(model, text_to_classify)
                if classifications is not None:
                    break
                time.sleep(rate_limit_pause)
            else:
                print(f"Error classifying batch starting at index {start}. Skipping...")
                continue

        try:
            # Update the DataFrame with the classifications
            for i, row in enumerate(classifications):
                for j, col in enumerate(category_columns):
                    df.at[start + i, col] = int(row[j])
        except Exception as e:
            for _ in range(2):
                classifications = classify_text(model, text_to_classify)
                if classifications is not None:
                    break
                time.sleep(rate_limit_pause)
            else:
                print(f"Error classifying batch starting at index {start}. Skipping...")
                continue
        
        try:
            for i, row in enumerate(classifications):
                for j, col in enumerate(category_columns):
                    df.at[start + i, col] = int(row[j])
        except Exception as e:
            print(f"Error updating DataFrame: {e}")
            continue
        
        time.sleep(rate_limit_pause)
    
    # Apply word segmentation to the content column
    df['content'] = df['content'].apply(lambda x: word_segmentation_vi(str(x)))
    
    # Save processed file, export columns of category_columns is int
    for col in category_columns:
        df[col] = df[col].astype(int)
    # Drop label_id column if it exists
    if 'label_id' in df.columns:
        df.drop(columns=['label_id'], inplace=True)
    df.to_csv(output_file, index=False)
    print(f"Saved processed file to {output_file}")

def main():
    parser = argparse.ArgumentParser(description="Process ViHSD CSV files with Google Generative AI")
    parser.add_argument("--input_dir", required=True, help="Directory containing input CSV files")
    parser.add_argument("--output_dir", required=True, help="Directory to save processed files")
    parser.add_argument("--api_key", required=True, help="Google Generative AI API key")
    parser.add_argument("--pause", type=float, default=4.0, help="Pause between API calls (seconds)")
    
    args = parser.parse_args()
    
    # Ensure output directory exists
    os.makedirs(args.output_dir, exist_ok=True)
    
    # Setup Google Generative AI
    model = setup_genai(args.api_key)
    
    # Get all CSV files in the input directory
    csv_files = glob.glob(os.path.join(args.input_dir, "*.csv"))
    if not csv_files:
        print(f"No CSV files found in {args.input_dir}")
        return
        
    print(f"Found {len(csv_files)} CSV files to process")
    
    # Process each file
    for input_file in csv_files:
        output_file = os.path.join(args.output_dir, os.path.basename(input_file))
        if os.path.exists(output_file):
            print(f"Output file {output_file} already exists. Skipping...")
            continue
        process_file(input_file, output_file, model, args.pause)

if __name__ == "__main__":
    # This script is used to process ViHSD CSV files with Google Generative AI
    # First, git clone from https://huggingface.co/datasets/sonlam1102/vihsd
    main()