kannada-gemma-4b / prepare_sft_data.py
harshaperla's picture
script fix
3f886b9
Raw
History Blame Contribute Delete
10.2 kB
#!/usr/bin/env python3
"""
Kannada SFT Data Preparation Script
====================================
Prepares instruction-tuning datasets for Kannada SFT.
Datasets used:
1. Cognitive-Lab/Kannada_Bilingual_Instruct (instruction, output, translation)
2. ai4bharat/indic-align (multiple configs: Dolly, WikiHow, etc.)
3. Cognitive-Lab/hh_dpo_kannada_translated (converted to SFT format from chosen responses)
Output: ./data/sft_kannada/ (parquet shards in conversational messages format)
The SFTTrainer expects either:
- {"text": "..."} for language modeling
- {"messages": [{"role": "user", "content": "..."}, {"role": "assistant", "content": "..."}]}
- {"prompt": "...", "completion": "..."}
We produce the conversational messages format.
Usage:
python prepare_sft_data.py [--output_dir ./data/sft_kannada]
"""
import argparse
import os
import sys
import logging
import json
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
log = logging.getLogger(__name__)
# =============================================================================
# Dataset loaders — each returns list of {"messages": [...]}
# =============================================================================
def load_kannada_bilingual_instruct():
"""
Cognitive-Lab/Kannada_Bilingual_Instruct
Columns: instruction, output, translation, text
Each row has Kannada instruction + Kannada output + English translation.
"""
from datasets import load_dataset
log.info("Loading Cognitive-Lab/Kannada_Bilingual_Instruct...")
try:
ds = load_dataset("Cognitive-Lab/Kannada_Bilingual_Instruct", split="train", streaming=True)
count = 0
for row in ds:
instruction = row.get("instruction", "").strip()
output = row.get("output", "").strip()
if not instruction or not output:
continue
yield {
"messages": [
{"role": "user", "content": instruction},
{"role": "assistant", "content": output},
],
"source": "kannada_bilingual_instruct",
}
count += 1
log.info(f" kannada_bilingual_instruct: {count:,} examples")
except Exception as e:
log.error(f" Failed: {e}")
def load_indic_align():
"""
ai4bharat/indic-align — multi-config dataset with translated instruction data.
Configs include: dolly_translated, wikihow_translated, vicuna_translated, etc.
We try to load Kannada-specific configs.
"""
from datasets import load_dataset
# These configs may have Kannada translations
configs_to_try = [
"kn", # direct Kannada config if available
]
log.info("Loading ai4bharat/indic-align...")
for config in configs_to_try:
try:
ds = load_dataset("ai4bharat/indic-align", config, split="train", streaming=True)
count = 0
for row in ds:
# indic-align typically has: instruction, input, output, language
instruction = row.get("instruction", "").strip()
inp = row.get("input", "").strip()
output = row.get("output", "").strip()
if inp:
instruction = f"{instruction}\n\n{inp}"
if not instruction or not output:
continue
yield {
"messages": [
{"role": "user", "content": instruction},
{"role": "assistant", "content": output},
],
"source": "indic_align",
}
count += 1
log.info(f" indic_align ({config}): {count:,} examples")
except Exception as e:
log.warning(f" indic-align config '{config}' failed: {e}")
# Also try loading the multi-language version and filter for Kannada
try:
ds = load_dataset("ai4bharat/indic-align", split="train", streaming=True)
count = 0
for row in ds:
lang = row.get("language", "").lower()
if "kan" not in lang and "kannada" not in lang:
continue
instruction = row.get("instruction", "").strip()
inp = row.get("input", "").strip()
output = row.get("output", "").strip()
if inp:
instruction = f"{instruction}\n\n{inp}"
if not instruction or not output:
continue
yield {
"messages": [
{"role": "user", "content": instruction},
{"role": "assistant", "content": output},
],
"source": "indic_align_filtered",
}
count += 1
log.info(f" indic_align (filtered kn): {count:,} examples")
except Exception as e:
log.warning(f" indic-align multi-language filtering failed: {e}")
def load_hh_dpo_kannada_as_sft():
"""
Cognitive-Lab/hh_dpo_kannada_translated
This is a DPO dataset (prompt, chosen, rejected) — we extract the chosen responses
as SFT data.
"""
from datasets import load_dataset
log.info("Loading Cognitive-Lab/hh_dpo_kannada_translated (as SFT)...")
try:
ds = load_dataset("Cognitive-Lab/hh_dpo_kannada_translated", split="train", streaming=True)
count = 0
for row in ds:
# DPO format: prompt, chosen, rejected
prompt = row.get("prompt", "").strip()
chosen = row.get("chosen", "").strip()
if not prompt or not chosen:
continue
yield {
"messages": [
{"role": "user", "content": prompt},
{"role": "assistant", "content": chosen},
],
"source": "hh_dpo_kannada_chosen",
}
count += 1
log.info(f" hh_dpo_kannada_chosen: {count:,} examples")
except Exception as e:
log.error(f" Failed: {e}")
def load_kannada_instruct_dataset():
"""
Cognitive-Lab/Kannada-Instruct-dataset (may be gated)
"""
from datasets import load_dataset
log.info("Loading Cognitive-Lab/Kannada-Instruct-dataset...")
try:
ds = load_dataset("Cognitive-Lab/Kannada-Instruct-dataset", split="train", streaming=True)
count = 0
for row in ds:
instruction = row.get("instruction", row.get("input", "")).strip()
output = row.get("output", row.get("response", "")).strip()
if not instruction or not output:
continue
yield {
"messages": [
{"role": "user", "content": instruction},
{"role": "assistant", "content": output},
],
"source": "kannada_instruct",
}
count += 1
log.info(f" kannada_instruct: {count:,} examples")
except Exception as e:
log.warning(f" Kannada-Instruct-dataset failed (may be gated): {e}")
def load_aya_kannada():
"""
Cognitive-Lab/Aya_Kannada (may be gated)
Part of the Cohere Aya collection.
"""
from datasets import load_dataset
log.info("Loading Cognitive-Lab/Aya_Kannada...")
try:
ds = load_dataset("Cognitive-Lab/Aya_Kannada", "aya_dataset", split="train", streaming=True)
count = 0
for row in ds:
# Aya format varies; try common column names
instruction = row.get("instruction", row.get("inputs", "")).strip()
output = row.get("output", row.get("targets", "")).strip()
if not instruction or not output:
continue
yield {
"messages": [
{"role": "user", "content": instruction},
{"role": "assistant", "content": output},
],
"source": "aya_kannada",
}
count += 1
log.info(f" aya_kannada: {count:,} examples")
except Exception as e:
log.warning(f" Aya_Kannada failed (may be gated): {e}")
# =============================================================================
# Main
# =============================================================================
def main():
parser = argparse.ArgumentParser(description="Prepare Kannada SFT data")
parser.add_argument("--output_dir", default="./data/sft_kannada",
help="Output directory")
parser.add_argument("--shard_size", type=int, default=50000,
help="Examples per shard")
args = parser.parse_args()
os.makedirs(args.output_dir, exist_ok=True)
log.info("=" * 70)
log.info("Kannada SFT Data Preparation")
log.info("=" * 70)
all_loaders = [
load_kannada_bilingual_instruct,
load_indic_align,
load_hh_dpo_kannada_as_sft,
load_kannada_instruct_dataset,
load_aya_kannada,
]
import pyarrow as pa
import pyarrow.parquet as pq
shard_num = 0
current_shard = []
total = 0
for loader in all_loaders:
for item in loader():
current_shard.append(item)
total += 1
if len(current_shard) >= args.shard_size:
out_path = os.path.join(args.output_dir, f"sft_shard_{shard_num:05d}.parquet")
table = pa.Table.from_pylist(current_shard)
pq.write_table(table, out_path)
log.info(f"Wrote {out_path}: {len(current_shard):,} examples")
shard_num += 1
current_shard = []
# Write remaining
if current_shard:
out_path = os.path.join(args.output_dir, f"sft_shard_{shard_num:05d}.parquet")
table = pa.Table.from_pylist(current_shard)
pq.write_table(table, out_path)
log.info(f"Wrote {out_path}: {len(current_shard):,} examples")
log.info("=" * 70)
log.info(f"DONE: {total:,} examples in {shard_num + 1} shards")
log.info(f"Output: {args.output_dir}")
log.info("=" * 70)
if __name__ == "__main__":
main()