Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 16 additions & 10 deletions dataflow/operators/general_text/eval/ngram_sample_evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,12 +7,17 @@
from dataflow import get_logger
from typing import Literal

_HAN_CHARACTER_PATTERN = re.compile(
r"[\u3400-\u4dbf\u4e00-\u9fff\uf900-\ufaff"
r"\U00020000-\U0002fa1f\U00030000-\U000323af]"
)

@OPERATOR_REGISTRY.register()
class NgramSampleEvaluator(OperatorABC):

def __init__(self, ngrams: int = 5, language: Literal['zh', 'en'] = 'en'):
if language not in ['zh', 'en']:
raise ValueError(f"Unsupported language: '{language}'. Supported options are: ['zh', 'en'].")
def __init__(self, ngrams: int = 5, language: Literal['zh', 'en', 'auto'] = 'en'):
if language not in ['zh', 'en', 'auto']:
raise ValueError(f"Unsupported language: '{language}'. Supported options are: ['zh', 'en', 'auto'].")
self.logger = get_logger()
self.logger.info(f'Initializing {self.__class__.__name__}...')
self.ngrams = ngrams
Expand All @@ -29,7 +34,7 @@ def get_desc(lang: str = "en"):
"支持中文(字级别)和英文(词级别)模式。\n"
"初始化参数:\n"
"- ngrams: n-gram长度,默认为5\n"
"- language: 处理语言'zh' 使用字粒度切分,其他使用空格分词,默认为 'en'\n"
"- language: 处理语言'zh' 使用字粒度切分,'en' 使用空格分词,'auto' 根据每条文本是否包含汉字自动选择,默认为 'en'\n"
"输出参数:\n"
"- NgramScore: n-gram重复比例得分(0到1之间,得分越高表示重复比例越低)"
)
Expand All @@ -39,7 +44,7 @@ def get_desc(lang: str = "en"):
"Supports Chinese (character-level) and English (word-level) modes.\n\n"
"Initialization Parameters:\n"
"- ngrams: Length of n-grams, default is 5.\n"
"- language: Processing language. 'zh' for character-level splitting, others for whitespace splitting. Default is 'en'.\n\n"
"- language: Processing language. 'zh' uses character-level splitting, 'en' uses whitespace splitting, and 'auto' selects per sample based on the presence of Han characters. Default is 'en'.\n\n"
"Output Parameters:\n"
"- NgramScore: N-gram repetition ratio score (0-1, higher score means less repetition/higher originality)."
)
Expand All @@ -53,8 +58,11 @@ def _score_func(self, sample):
# 移除标点符号
content = re.sub(r'[^\w\s]', '', content)

# --- 根据语言选择切分逻辑 ---
if self.language == 'zh':
# Auto detection is opt-in so an explicit language mode remains predictable.
use_character_tokens = self.language == 'zh' or (
self.language == 'auto' and _HAN_CHARACTER_PATTERN.search(content)
)
if use_character_tokens:
# 中文模式:去除所有空格,按“字”切分
content = re.sub(r'\s+', '', content)
tokens = list(content)
Expand All @@ -63,8 +71,6 @@ def _score_func(self, sample):
# 默认/英文模式:按“空格”切分
tokens = content.split()
join_char = " "
# ---------------------------

if len(tokens) < self.ngrams:
return 0.0

Expand All @@ -90,4 +96,4 @@ def run(self, storage: DataFlowStorage, input_key: str, output_key: str='NgramSc
dataframe = storage.read("dataframe")
scores = self.eval(dataframe, input_key)
dataframe[self.output_key] = scores
storage.write(dataframe)
storage.write(dataframe)
12 changes: 6 additions & 6 deletions dataflow/operators/general_text/filter/ngram_filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,9 +8,9 @@
@OPERATOR_REGISTRY.register()
class NgramFilter(OperatorABC):

def __init__(self, min_score=0.8, max_score=1, ngrams=5, language: Literal['zh', 'en'] = 'en'):
if language not in ['zh', 'en']:
raise ValueError(f"Unsupported language: '{language}'. Supported options are: ['zh', 'en'].")
def __init__(self, min_score=0.8, max_score=1, ngrams=5, language: Literal['zh', 'en', 'auto'] = 'en'):
if language not in ['zh', 'en', 'auto']:
raise ValueError(f"Unsupported language: '{language}'. Supported options are: ['zh', 'en', 'auto'].")
self.logger = get_logger()
self.min_score = min_score
self.max_score = max_score
Expand All @@ -26,6 +26,7 @@ def get_desc(lang: str = "zh"):
"- min_score:最小n-gram得分阈值\n"
"- max_score:最大n-gram得分阈值\n"
"- ngrams:n-gram的n值\n"
"- language:处理语言;'zh' 使用字粒度切分,'en' 使用空格分词,'auto' 根据每条文本是否包含汉字自动选择\n"
"输出参数:\n"
"- 过滤后的DataFrame,仅保留n-gram得分在指定范围内的文本\n"
"- 返回包含n-gram得分字段名的列表"
Expand All @@ -36,7 +37,8 @@ def get_desc(lang: str = "zh"):
"Input Parameters:\n"
"- min_score: Minimum n-gram score threshold\n"
"- max_score: Maximum n-gram score threshold\n"
"- ngrams: n value for n-gram\n\n"
"- ngrams: n value for n-gram\n"
"- language: Processing language. 'zh' uses character-level splitting, 'en' uses whitespace splitting, and 'auto' selects per sample based on the presence of Han characters.\n\n"
"Output Parameters:\n"
"- Filtered DataFrame containing only texts with n-gram score within specified range\n"
"- List containing n-gram score field name"
Expand All @@ -53,5 +55,3 @@ def run(self, storage: DataFlowStorage, input_key: str, output_key: str='NgramSc
output_file = storage.write(filtered_dataframe)
self.logger.info(f"Filtering completed. Total records passing filter: {len(filtered_dataframe)}.")
return [self.output_key]


44 changes: 44 additions & 0 deletions test/cpu_only/test_ngram_sample_evaluator.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,44 @@
import pytest

from dataflow.operators.general_text.eval.ngram_sample_evaluator import (
NgramSampleEvaluator,
)


@pytest.mark.cpu
def test_auto_mode_uses_character_ngrams_for_han_text():
evaluator = NgramSampleEvaluator(ngrams=5, language="auto")

assert evaluator._score_func("今天天气真不错,适合出门散步。") == 1.0


@pytest.mark.cpu
def test_auto_mode_keeps_word_ngrams_for_english_text():
text = "test test test test test test final"
auto_evaluator = NgramSampleEvaluator(ngrams=5, language="auto")
en_evaluator = NgramSampleEvaluator(ngrams=5, language="en")

assert auto_evaluator._score_func(text) == en_evaluator._score_func(text)


@pytest.mark.cpu
def test_default_and_explicit_en_modes_are_not_overridden_for_mixed_text():
default_evaluator = NgramSampleEvaluator(ngrams=5)
en_evaluator = NgramSampleEvaluator(ngrams=5, language="en")
text = "test test test test test test 中文"

assert default_evaluator._score_func(text) == pytest.approx(2 / 3)
assert en_evaluator._score_func(text) == pytest.approx(2 / 3)


@pytest.mark.cpu
def test_auto_mode_detects_han_characters_outside_basic_block():
evaluator = NgramSampleEvaluator(ngrams=5, language="auto")

assert evaluator._score_func("𠀀𠀁𠀂𠀃𠀄𠀅") == 1.0


@pytest.mark.cpu
def test_rejects_unsupported_language():
with pytest.raises(ValueError, match="Unsupported language"):
NgramSampleEvaluator(language="cjk")
Loading