File size: 2,125 Bytes
cfd7c14
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

import math
import re
import string
from typing import Any, Iterable, List

import numpy as np

try:
    import pandas as pd
except ImportError:  # pragma: no cover
    pd = None


def is_missing(value: Any) -> bool:
    if value is None:
        return True
    if isinstance(value, float) and math.isnan(value):
        return True
    if pd is not None:
        try:
            missing = pd.isna(value)
            if isinstance(missing, (bool, np.bool_)):
                return bool(missing)
        except (TypeError, ValueError):
            pass
    return False


def coerce_text(value: Any) -> str:
    if is_missing(value):
        return ""
    return str(value)


def extract_chat_prompt(value: Any) -> str:
    if is_missing(value):
        return ""
    if isinstance(value, str):
        return value
    if isinstance(value, np.ndarray):
        value = value.tolist()
    if isinstance(value, (list, tuple)):
        if not value:
            return ""
        first = value[0]
        if isinstance(first, dict):
            return coerce_text(first.get("content", ""))
        return coerce_text(first)
    if isinstance(value, dict):
        return coerce_text(value.get("content", ""))
    return coerce_text(value)


def extract_prompt_from_row(
    row: Any,
    prompt_column: str = "prompt",
    legacy_input_column: str = "input",
) -> str:
    if prompt_column in row and not is_missing(row[prompt_column]):
        return coerce_text(row[prompt_column])
    if legacy_input_column in row:
        return extract_chat_prompt(row[legacy_input_column])
    return ""


def preprocess_text(text: Any) -> str:
    text = coerce_text(text).lower()
    text = re.sub(r"\d+", " ", text)
    common_punct = string.punctuation.replace("-", "")
    text = text.translate(str.maketrans(common_punct, " " * len(common_punct)))
    text = re.sub(r"\s+", " ", text)
    return text.strip()


def preprocess_batch(texts: Iterable[Any]) -> List[str]:
    return [preprocess_text(text) for text in texts]


def simple_tokenizer(text: str) -> List[str]:
    return text.split()