chatbot / make_data.py
ogx786's picture
Update make_data.py
b5a7e08 verified
Raw
History Blame Contribute Delete
2.58 kB
import pandas as pd
from sklearn.model_selection import train_test_split
import csv
import os
INPUT_FILE = "hbl_nadra_clean.csv"
OUTPUT_DIR = "splits"
os.makedirs(OUTPUT_DIR, exist_ok=True)
# ============================================================
# LOAD DATA
# ============================================================
print("Loading dataset...")
df = pd.read_csv(
INPUT_FILE,
encoding="utf-8"
)
print("\nOriginal columns:")
print(df.columns)
print("\nTotal rows:", len(df))
# ============================================================
# BASIC VALIDATION
# ============================================================
df = df.dropna()
df["urdu"] = df["urdu"].astype(str).str.strip()
df["roman"] = df["roman"].astype(str).str.strip()
# remove empty rows
df = df[
(df["urdu"] != "") &
(df["roman"] != "")
]
print("After cleaning:", len(df))
# ============================================================
# TRAIN / VAL / TEST SPLIT
# 90 / 5 / 5
# ============================================================
train_df, temp_df = train_test_split(
df,
test_size=0.10,
random_state=42
)
val_df, test_df = train_test_split(
temp_df,
test_size=0.50,
random_state=42
)
print("\nSplit sizes")
print("----------------")
print("Train:", len(train_df))
print("Validation:", len(val_df))
print("Test:", len(test_df))
# ============================================================
# SAVE FUNCTION
# ============================================================
def save_csv(data, filename):
path = os.path.join(
OUTPUT_DIR,
filename
)
data.to_csv(
path,
index=False,
encoding="utf-8-sig",
quoting=csv.QUOTE_ALL
)
print("Saved:", path)
# ============================================================
# SAVE SPLITS
# ============================================================
save_csv(train_df, "train.csv")
save_csv(val_df, "val.csv")
save_csv(test_df, "test.csv")
# ============================================================
# VERIFY
# ============================================================
print("\n\nChecking saved files\n")
for file in [
"train.csv",
"val.csv",
"test.csv"
]:
path = os.path.join(
OUTPUT_DIR,
file
)
check = pd.read_csv(
path,
encoding="utf-8-sig"
)
print("\n========================")
print(file)
print("========================")
print(check.head(3).to_string())
print("\nColumns:")
print(check.columns.tolist())