Spaces:
Sleeping
Sleeping
File size: 1,759 Bytes
33480f6 ef8438b 33480f6 14b08d6 3c64996 33480f6 ef8438b 33480f6 14b08d6 33480f6 ef8438b 14b08d6 | 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 | import polars as pl
from iterstrat.ml_stratifiers import MultilabelStratifiedShuffleSplit
from langchain_text_splitters import RecursiveCharacterTextSplitter
import numpy as np
from datasets import Dataset
from scripts.eda import training_set, class_cols
WORD_LIMIT=512
CHAR_LIMIT = WORD_LIMIT * 6
CHUNK_OVERLAP = 0
splitter = RecursiveCharacterTextSplitter(
chunk_size=CHAR_LIMIT,
chunk_overlap=CHUNK_OVERLAP,
separators=["\n\n", "\n", ". ", " ", ""],
)
#trying this kind of coding haha
def word_count(text: str) -> int:
return len(text.split())
def split_row(row: dict) -> list[dict]:
chunks = splitter.split_text(row["comment_text"])
result = []
for n, chunk in enumerate(chunks):
new_row = row.copy()
new_row["comment_text"] = chunk
new_row["id"] = f"{row["id"]}__chunk{n}"
result.append(new_row)
return result
def process_df(df:pl.DataFrame) -> pl.DataFrame:
kept_rows: list[dict] = []
expanded_rows: list[dict] = []
for row in df.iter_rows(named=True):
if word_count(row["comment_text"]) > WORD_LIMIT:
expanded_rows.extend(split_row(row))
else:
kept_rows.append(row)
all_rows = kept_rows + expanded_rows
return pl.DataFrame(all_rows, schema=df.schema)
bert_training_set = process_df(training_set)
y = bert_training_set.select(class_cols).to_numpy()
X_dummy = np.zeros((bert_training_set.height, 1))
msss = MultilabelStratifiedShuffleSplit(
n_splits=1, test_size=0.1, random_state=42
)
train_idx, val_idx = next(msss.split(X_dummy, y))
train_set = bert_training_set[train_idx]
val_set = bert_training_set[val_idx]
print(train_set.height)
train_ds = Dataset.from_polars(train_set)
val_ds = Dataset.from_polars(val_set)
|