FannyFa-Model-V1 / data_loader.py
FannyFa's picture
Upload 16 files
3eecd6b verified
Raw History Blame Contribute Delete
31 kB
import os
import json
import csv
import re
import hashlib
import statistics
import time # FIXED: Missing import
from pathlib import Path
from typing import List, Dict, Tuple, Optional, Any, Iterator, Set
from collections import Counter
from dataclasses import dataclass, field
from config import DATA_DIR, DEFAULT_MEMORY_LIMIT_GB
from utils import console, Theme, debug_logger, error_logger
@dataclass
class DatasetStats:
total_samples: int = 0
valid_samples: int = 0
invalid_samples: int = 0
duplicate_samples: int = 0
avg_length: float = 0.0
max_length: int = 0
min_length: int = 0
avg_words: float = 0.0
std_length: float = 0.0
format_detected: str = ""
structure_type: str = ""
language: str = "unknown"
has_headers: bool = False
column_names: List[str] = field(default_factory=list)
warnings: List[str] = field(default_factory=list)
conversation_pairs: List[Tuple[str, str]] = field(default_factory=list)
word_frequency: Dict[str, int] = field(default_factory=dict)
char_frequency: Dict[str, int] = field(default_factory=dict)
def to_dict(self) -> Dict:
import dataclasses
return dataclasses.asdict(self)
class EnhancedDatasetLoader:
def __init__(
self, chunk_size: int = 10000, memory_limit_gb: float = DEFAULT_MEMORY_LIMIT_GB
):
self.chunk_size = chunk_size
self.memory_limit_gb = memory_limit_gb
self.stats = DatasetStats()
self._cache: Dict[str, Tuple[List[str], DatasetStats]] = {}
self._loaded_files: Set[str] = set()
self._resume_state: Dict[str, int] = {}
def load(
self, filepath: str, augment: bool = False, augment_cfg: Dict = None
) -> Tuple[List[str], DatasetStats]:
if not os.path.exists(filepath):
raise FileNotFoundError(f"Dataset not found: {filepath}")
if os.path.getsize(filepath) == 0:
raise ValueError("Dataset file is empty")
file_hash = self._get_file_hash(filepath)
if file_hash in self._cache:
console.print(Theme.dim("Using cached dataset"))
return self._cache[file_hash]
ext = Path(filepath).suffix.lower()
self.stats.format_detected = ext[1:].upper() if ext else "UNKNOWN"
try:
loader_map = {
".txt": self._load_txt,
".json": self._load_json,
".jsonl": self._load_jsonl,
".csv": self._load_csv,
".tsv": self._load_tsv,
".parquet": self._load_parquet,
".arrow": self._load_arrow,
".json.gz": self._load_compressed_json,
".jsonl.gz": self._load_compressed_jsonl,
}
if ext not in loader_map:
raise ValueError(f"Unsupported format: {ext}")
samples = loader_map[ext](filepath)
samples = self._validate_samples(samples)
if samples:
self.stats.language = self._detect_language(samples[:100])
self._compute_stats(samples)
if augment and augment_cfg:
samples = self._augment_samples(samples, augment_cfg)
self._cache[file_hash] = (samples, self.stats)
self._loaded_files.add(filepath)
return samples, self.stats
except Exception as e:
error_logger.error(f"Error loading dataset {filepath}: {str(e)}")
raise
def load_streaming_with_resume(
self, filepath: str, checkpoint_file: str = None
) -> Iterator[List[str]]:
if not os.path.exists(filepath):
raise FileNotFoundError(f"Dataset not found: {filepath}")
ext = Path(filepath).suffix.lower()
resume_pos = 0
if checkpoint_file and os.path.exists(checkpoint_file):
try:
with open(checkpoint_file, "r") as f:
state = json.load(f)
resume_pos = state.get("position", 0)
console.print(Theme.dim(f"Resuming from position {resume_pos}"))
except Exception as e:
console.print(Theme.warning(f"Could not load resume state: {e}"))
# ------------------------------------------------------------------
# FIX: lazy lookup via string name — supaya AttributeError hanya
# muncul kalau method-nya benar-benar dipanggil, bukan saat dict
# dibuat (sebelumnya semua ekstensi crash karena salah satu method
# tidak ada).
# ------------------------------------------------------------------
stream_methods = {
".txt": "_stream_txt_with_resume",
".jsonl": "_stream_jsonl_with_resume",
".csv": "_stream_csv_with_resume",
".tsv": "_stream_tsv_with_resume",
".parquet": "_stream_parquet_with_resume",
}
if ext in stream_methods:
method = getattr(self, stream_methods[ext], None)
if method is None:
debug_logger.debug(
f"Streaming method untuk {ext} tidak tersedia, fallback ke load()"
)
samples, _ = self.load(filepath)
yield samples
else:
yield from method(filepath, resume_pos, checkpoint_file)
else:
samples, _ = self.load(filepath)
yield samples
def _save_resume_state(self, checkpoint_file: str, position: int) -> None:
"""
Simpan posisi resume. FIX: sebelumnya ada orphan code yang
mereferensikan variabel `content` yang tidak ada di scope.
"""
try:
with open(checkpoint_file, "w") as f:
json.dump({"position": position, "timestamp": time.time()}, f)
except Exception as e:
debug_logger.debug(f"Could not save resume state: {e}")
def _is_conversation_format(self, content: str) -> bool:
lines = content.split("\n")
markers = [
"User:",
"AI:",
"Assistant:",
"Human:",
"Bot:",
"You:",
"System:",
"Agent:",
]
return (
sum(1 for line in lines[:30] if any(line.startswith(m) for m in markers))
> 3
)
def _parse_conversation_format(self, content: str) -> List[str]:
samples = []
current_conv = []
for line in content.split("\n"):
line = line.strip()
if not line:
if current_conv:
samples.append("\n".join(current_conv))
current_conv = []
else:
current_conv.append(line)
if current_conv:
samples.append("\n".join(current_conv))
return samples
def _load_jsonl(self, filepath: str) -> List[str]:
samples = []
invalid_count = 0
detected_format = "unknown"
with open(filepath, "r", encoding="utf-8") as f:
for idx, line in enumerate(f):
line = line.strip()
if not line:
continue
try:
obj = json.loads(line)
if not isinstance(obj, dict):
invalid_count += 1
continue
text_sample = self._extract_conversation_from_json(obj)
if text_sample:
samples.append(text_sample)
detected_format = self._detect_json_format(obj)
except json.JSONDecodeError:
invalid_count += 1
if invalid_count <= 5:
self.stats.warnings.append(f"Invalid JSON at line {idx + 1}")
except Exception as e:
invalid_count += 1
if invalid_count <= 5:
self.stats.warnings.append(
f"Error at line {idx + 1}: {str(e)[:50]}"
)
if invalid_count > 5:
self.stats.warnings.append(
f"... and {invalid_count - 5} more invalid lines"
)
self.stats.structure_type = detected_format
console.print(Theme.dim(f"Detected format: {detected_format}"))
return samples
def _extract_conversation_from_json(self, obj: Dict) -> Optional[str]:
if "user" in obj and "assistant" in obj:
user_text = self._clean_text_for_conversation(str(obj["user"]))
assistant_text = self._clean_text_for_conversation(str(obj["assistant"]))
if user_text and assistant_text:
self.stats.conversation_pairs.append((user_text, assistant_text))
return f"User: {user_text}\nAI: {assistant_text}"
if "prompt" in obj and "response" in obj:
prompt_text = self._clean_text_for_conversation(str(obj["prompt"]))
response_text = self._clean_text_for_conversation(str(obj["response"]))
if prompt_text and response_text:
self.stats.conversation_pairs.append((prompt_text, response_text))
return f"User: {prompt_text}\nAI: {response_text}"
if "instruction" in obj and "response" in obj:
inst_text = self._clean_text_for_conversation(str(obj["instruction"]))
resp_text = self._clean_text_for_conversation(str(obj["response"]))
if inst_text and resp_text:
self.stats.conversation_pairs.append((inst_text, resp_text))
return f"User: {inst_text}\nAI: {resp_text}"
if "text" in obj:
text = self._clean_text_for_conversation(str(obj["text"]))
if text:
return text
if "messages" in obj and isinstance(obj["messages"], list):
messages = obj["messages"]
conv_parts = []
for msg in messages:
if isinstance(msg, dict):
role = msg.get("role", "")
content = msg.get("content", "")
if role and content:
role_map = {
"user": "User",
"assistant": "AI",
"system": "System",
}
role_display = role_map.get(role, role.capitalize())
conv_parts.append(f"{role_display}: {content}")
if conv_parts:
return "\n".join(conv_parts)
if "question" in obj and "answer" in obj:
q_text = self._clean_text_for_conversation(str(obj["question"]))
a_text = self._clean_text_for_conversation(str(obj["answer"]))
if q_text and a_text:
self.stats.conversation_pairs.append((q_text, a_text))
return f"User: {q_text}\nAI: {a_text}"
if "input" in obj and "output" in obj:
in_text = self._clean_text_for_conversation(str(obj["input"]))
out_text = self._clean_text_for_conversation(str(obj["output"]))
if in_text and out_text:
self.stats.conversation_pairs.append((in_text, out_text))
return f"User: {in_text}\nAI: {out_text}"
extracted = self._extract_text_from_obj_fallback(obj)
if extracted:
return extracted
return None
def _detect_json_format(self, obj: Dict) -> str:
if "user" in obj and "assistant" in obj:
return "user_assistant"
if "prompt" in obj and "response" in obj:
return "prompt_response"
if "instruction" in obj and "response" in obj:
return "instruction_response"
if "messages" in obj:
return "messages"
if "question" in obj and "answer" in obj:
return "qa"
if "text" in obj:
return "text"
return "unknown"
def _extract_text_from_obj_fallback(self, obj: Dict) -> Optional[str]:
for key in ["content", "data", "value", "description"]:
if key in obj and isinstance(obj[key], str):
text = self._clean_text_for_conversation(obj[key])
if text:
return text
return None
def _load_txt(self, filepath: str) -> List[str]:
samples = []
with open(filepath, "r", encoding="utf-8") as f:
current_sample = ""
for line in f:
line = line.rstrip()
if line.strip():
current_sample += line + " "
elif current_sample:
samples.append(current_sample.strip())
current_sample = ""
if current_sample:
samples.append(current_sample.strip())
return samples
def _load_json(self, filepath: str) -> List[str]:
samples = []
try:
with open(filepath, "r", encoding="utf-8") as f:
data = json.load(f)
if isinstance(data, list):
for item in data:
if isinstance(item, dict):
text = self._extract_conversation_from_json(item)
if text:
samples.append(text)
elif isinstance(item, str):
if item.strip():
samples.append(item)
elif isinstance(data, dict):
text = self._extract_conversation_from_json(data)
if text:
samples.append(text)
except json.JSONDecodeError as e:
error_logger.error(f"JSON decode error: {e}")
return samples
def _load_csv(self, filepath: str) -> List[str]:
samples = []
try:
with open(filepath, "r", encoding="utf-8") as f:
reader = csv.DictReader(f)
for row in reader:
text_parts = []
for col_name, col_lower in [(k, k.lower()) for k in row.keys()]:
if col_lower in [
"user",
"prompt",
"question",
"input",
"human",
"text",
]:
value = row[col_name]
if value and value.strip():
text_parts.append(value.strip())
elif col_lower in [
"assistant",
"response",
"answer",
"output",
"ai",
]:
value = row[col_name]
if value and value.strip():
text_parts.append(value.strip())
if text_parts:
samples.append(" ".join(text_parts))
except Exception as e:
error_logger.error(f"CSV load error: {e}")
return samples
def _load_tsv(self, filepath: str) -> List[str]:
return self._load_csv_like(filepath, delimiter="\t")
def _load_csv_like(self, filepath: str, delimiter: str = ",") -> List[str]:
samples = []
try:
with open(filepath, "r", encoding="utf-8") as f:
reader = csv.DictReader(f, delimiter=delimiter)
for row in reader:
text_parts = []
for col_name, col_lower in [(k, k.lower()) for k in row.keys()]:
if col_lower in [
"user",
"prompt",
"question",
"input",
"human",
"text",
]:
value = row[col_name]
if value and value.strip():
text_parts.append(value.strip())
elif col_lower in [
"assistant",
"response",
"answer",
"output",
"ai",
]:
value = row[col_name]
if value and value.strip():
text_parts.append(value.strip())
if text_parts:
samples.append(" ".join(text_parts))
except Exception as e:
error_logger.error(f"CSV-like load error: {e}")
return samples
def _load_parquet(self, filepath: str) -> List[str]:
try:
import pandas as pd
except ImportError:
raise ImportError("pandas required for parquet files")
samples = []
try:
df = pd.read_parquet(filepath)
for _, row in df.iterrows():
text_parts = []
for col in df.columns:
col_lower = str(col).lower()
if col_lower in [
"user",
"prompt",
"question",
"input",
"human",
"text",
]:
value = str(row[col])
if value and value.strip():
text_parts.append(value.strip())
elif col_lower in [
"assistant",
"response",
"answer",
"output",
"ai",
]:
value = str(row[col])
if value and value.strip():
text_parts.append(value.strip())
if text_parts:
samples.append(" ".join(text_parts))
except Exception as e:
error_logger.error(f"Parquet load error: {e}")
return samples
def _load_arrow(self, filepath: str) -> List[str]:
try:
import pyarrow.parquet as pq
except ImportError:
raise ImportError("pyarrow required for arrow files")
samples = []
try:
table = pq.read_table(filepath)
df = table.to_pandas()
for _, row in df.iterrows():
text_parts = []
for col in df.columns:
col_lower = str(col).lower()
if col_lower in [
"user",
"prompt",
"question",
"input",
"human",
"text",
]:
value = str(row[col])
if value and value.strip():
text_parts.append(value.strip())
elif col_lower in [
"assistant",
"response",
"answer",
"output",
"ai",
]:
value = str(row[col])
if value and value.strip():
text_parts.append(value.strip())
if text_parts:
samples.append(" ".join(text_parts))
except Exception as e:
error_logger.error(f"Arrow load error: {e}")
return samples
def _load_compressed_json(self, filepath: str) -> List[str]:
import gzip
samples = []
try:
with gzip.open(filepath, "rt", encoding="utf-8") as f:
data = json.load(f)
if isinstance(data, list):
for item in data:
if isinstance(item, dict):
text = self._extract_conversation_from_json(item)
if text:
samples.append(text)
except Exception as e:
error_logger.error(f"Compressed JSON load error: {e}")
return samples
def _load_compressed_jsonl(self, filepath: str) -> List[str]:
import gzip
samples = []
try:
with gzip.open(filepath, "rt", encoding="utf-8") as f:
for line in f:
line = line.strip()
if line:
try:
obj = json.loads(line)
if isinstance(obj, dict):
text = self._extract_conversation_from_json(obj)
if text:
samples.append(text)
except json.JSONDecodeError:
continue
except Exception as e:
error_logger.error(f"Compressed JSONL load error: {e}")
return samples
# ------------------------------------------------------------------
# Streaming methods (semua harus punya signature yang sama)
# ------------------------------------------------------------------
def _stream_txt_with_resume(
self, filepath: str, resume_pos: int, checkpoint_file: str
) -> Iterator[List[str]]:
chunk = []
current_pos = 0
with open(filepath, "r", encoding="utf-8") as f:
if resume_pos > 0:
for _ in range(resume_pos):
f.readline()
current_pos = resume_pos
for line in f:
current_pos += 1
line = line.strip()
if line:
chunk.append(line)
if len(chunk) >= self.chunk_size:
if checkpoint_file:
self._save_resume_state(checkpoint_file, current_pos)
yield chunk
chunk = []
if chunk and checkpoint_file:
self._save_resume_state(checkpoint_file, current_pos)
if chunk:
yield chunk
def _stream_jsonl_with_resume(
self, filepath: str, resume_pos: int, checkpoint_file: str
) -> Iterator[List[str]]:
"""
Stream JSONL dengan dukungan resume.
Method ini sebelumnya HILANG di kode asli — ditambahkan di sini.
"""
chunk: List[str] = []
current_pos = 0
with open(filepath, "r", encoding="utf-8") as f:
if resume_pos > 0:
for _ in range(resume_pos):
f.readline()
current_pos = resume_pos
for line in f:
current_pos += 1
line = line.strip()
if not line:
continue
try:
obj = json.loads(line)
except json.JSONDecodeError:
continue
if not isinstance(obj, dict):
continue
text = self._extract_conversation_from_json(obj)
if text:
chunk.append(text)
if len(chunk) >= self.chunk_size:
if checkpoint_file:
self._save_resume_state(checkpoint_file, current_pos)
yield chunk
chunk = []
if chunk and checkpoint_file:
self._save_resume_state(checkpoint_file, current_pos)
if chunk:
yield chunk
def _stream_csv_with_resume(
self, filepath: str, resume_pos: int, checkpoint_file: str
) -> Iterator[List[str]]:
chunk = []
current_pos = 0
with open(filepath, "r", encoding="utf-8") as f:
reader = csv.DictReader(f)
if resume_pos > 0:
for _ in range(resume_pos - 1):
next(reader, None)
current_pos = resume_pos
for row in reader:
current_pos += 1
values = [v.strip() for v in row.values() if v.strip()]
if values:
chunk.append(" ".join(values))
if len(chunk) >= self.chunk_size:
if checkpoint_file:
self._save_resume_state(checkpoint_file, current_pos)
yield chunk
chunk = []
if chunk and checkpoint_file:
self._save_resume_state(checkpoint_file, current_pos)
if chunk:
yield chunk
def _stream_tsv_with_resume(
self, filepath: str, resume_pos: int, checkpoint_file: str
) -> Iterator[List[str]]:
chunk = []
current_pos = 0
with open(filepath, "r", encoding="utf-8") as f:
reader = csv.DictReader(f, delimiter="\t")
if resume_pos > 0:
for _ in range(resume_pos - 1):
next(reader, None)
current_pos = resume_pos
for row in reader:
current_pos += 1
values = [v.strip() for v in row.values() if v.strip()]
if values:
chunk.append(" ".join(values))
if len(chunk) >= self.chunk_size:
if checkpoint_file:
self._save_resume_state(checkpoint_file, current_pos)
yield chunk
chunk = []
if chunk and checkpoint_file:
self._save_resume_state(checkpoint_file, current_pos)
if chunk:
yield chunk
def _stream_parquet_with_resume(
self, filepath: str, resume_pos: int, checkpoint_file: str
) -> Iterator[List[str]]:
try:
import pandas as pd
except ImportError:
raise ImportError("pandas required for parquet files")
current_pos = 0
for chunk_df in pd.read_parquet(filepath, chunksize=self.chunk_size):
current_pos += len(chunk_df)
if current_pos < resume_pos:
continue
chunk = [
str(x).strip() for x in chunk_df.iloc[:, 0].tolist() if str(x).strip()
]
if chunk:
if checkpoint_file:
self._save_resume_state(checkpoint_file, current_pos)
yield chunk
def _get_file_hash(self, filepath: str) -> str:
hasher = hashlib.sha256()
with open(filepath, "rb") as f:
for chunk in iter(lambda: f.read(65536), b""):
hasher.update(chunk)
return hasher.hexdigest()
def _detect_language(self, samples: List[str]) -> str:
indo_words = {
"yang",
"dan",
"di",
"ke",
"dari",
"ini",
"itu",
"untuk",
"dengan",
"adalah",
"pada",
"dalam",
"atas",
"oleh",
"sebagai",
"akan",
"karena",
"atau",
}
eng_words = {
"the",
"of",
"and",
"to",
"in",
"for",
"on",
"at",
"by",
"with",
"from",
"up",
"about",
"into",
"through",
"during",
"including",
}
text = " ".join(samples[:50]).lower()
tokens = set(re.findall(r"\w+", text))
id_count = len(tokens & indo_words)
en_count = len(tokens & eng_words)
if id_count > en_count:
return "indonesian"
elif en_count > id_count:
return "english"
else:
return "mixed"
def _validate_samples(self, samples: List[str]) -> List[str]:
valid_samples = []
invalid_count = 0
seen = set()
duplicates = 0
for sample in samples:
if not sample or not sample.strip():
invalid_count += 1
continue
sample = self._clean_text_for_conversation(sample)
if len(sample) < 3:
invalid_count += 1
continue
if sample in seen:
duplicates += 1
continue
valid_samples.append(sample)
seen.add(sample)
self.stats.invalid_samples = invalid_count
self.stats.duplicate_samples = duplicates
self.stats.total_samples = len(samples)
self.stats.valid_samples = len(valid_samples)
return valid_samples
def _augment_samples(self, samples: List[str], cfg: Dict) -> List[str]:
out = list(samples)
if cfg.get("split_long", False):
max_len = cfg.get("split_max_chars", 300)
for s in samples:
if len(s) > max_len:
parts = re.split(r"(?<=[.!?])\s+", s)
for i in range(0, len(parts), 2):
chunk = " ".join(parts[i : i + 2]).strip()
if chunk:
out.append(chunk)
seen = set()
uniq = []
for s in out:
if s not in seen and s.strip():
seen.add(s)
uniq.append(s)
return uniq
def _compute_stats(self, samples: List[str]) -> None:
if not samples:
return
lengths = [len(s) for s in samples]
word_counts = [len(s.split()) for s in samples]
self.stats.max_length = max(lengths)
self.stats.min_length = min(lengths)
self.stats.avg_length = sum(lengths) / len(lengths)
self.stats.avg_words = sum(word_counts) / len(word_counts)
self.stats.std_length = statistics.stdev(lengths) if len(lengths) > 1 else 0
word_freq = Counter()
for s in samples[:1000]:
words = re.findall(r"\w+", s.lower())
word_freq.update(words)
self.stats.word_frequency = dict(word_freq.most_common(50))
char_freq = Counter()
for s in samples[:1000]:
char_freq.update(s.lower())
self.stats.char_frequency = dict(char_freq.most_common(30))
def _clean_text_for_conversation(self, text: str) -> str:
if not text:
return ""
text = str(text).strip()
text = re.sub(r"\s+", " ", text)
text = re.sub(r"[\n\r\t]+", " ", text)
return text