ru-stylometry 0.1.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- ru_stylometry/__init__.py +34 -0
- ru_stylometry/__main__.py +5 -0
- ru_stylometry/_lexicons.py +76 -0
- ru_stylometry/_morph.py +41 -0
- ru_stylometry/_stats.py +91 -0
- ru_stylometry/_text.py +99 -0
- ru_stylometry/cli.py +269 -0
- ru_stylometry/document.py +66 -0
- ru_stylometry/evaluation.py +267 -0
- ru_stylometry/explain.py +134 -0
- ru_stylometry/features/__init__.py +130 -0
- ru_stylometry/features/char.py +54 -0
- ru_stylometry/features/formatting.py +34 -0
- ru_stylometry/features/length.py +30 -0
- ru_stylometry/features/lexical.py +37 -0
- ru_stylometry/features/morph.py +109 -0
- ru_stylometry/features/readability.py +37 -0
- ru_stylometry/features/repetition.py +43 -0
- ru_stylometry/features/syntax.py +64 -0
- ru_stylometry/features/typography.py +57 -0
- ru_stylometry/model.py +196 -0
- ru_stylometry/perturb.py +116 -0
- ru_stylometry/py.typed +1 -0
- ru_stylometry/vectorizer.py +91 -0
- ru_stylometry-0.1.1.dist-info/METADATA +274 -0
- ru_stylometry-0.1.1.dist-info/RECORD +30 -0
- ru_stylometry-0.1.1.dist-info/WHEEL +5 -0
- ru_stylometry-0.1.1.dist-info/entry_points.txt +2 -0
- ru_stylometry-0.1.1.dist-info/licenses/LICENSE +21 -0
- ru_stylometry-0.1.1.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
"""ru-stylometry — interpretable stylometric features for Russian text.
|
|
2
|
+
|
|
3
|
+
Public API:
|
|
4
|
+
- :func:`extract_features` — named stylometric features of one text.
|
|
5
|
+
- :func:`feature_names` and :func:`describe_features` — the feature catalogue.
|
|
6
|
+
- :class:`StylometricVectorizer` — scikit-learn transformer: texts -> feature matrix.
|
|
7
|
+
- :class:`StylometricClassifier` — classifier over those features with explanations.
|
|
8
|
+
|
|
9
|
+
Features are grouped by the level of analysis (characters, typography, vocabulary, syntax,
|
|
10
|
+
readability, repetition, morphology); see :data:`ALL_GROUPS` and :data:`DEFAULT_GROUPS`.
|
|
11
|
+
The submodules ``evaluation``, ``explain`` and ``perturb`` hold the quality criteria, the
|
|
12
|
+
explanations and the text transformations used in robustness checks.
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
from ru_stylometry.features import (
|
|
16
|
+
ALL_GROUPS,
|
|
17
|
+
DEFAULT_GROUPS,
|
|
18
|
+
describe_features,
|
|
19
|
+
extract_features,
|
|
20
|
+
feature_names,
|
|
21
|
+
)
|
|
22
|
+
from ru_stylometry.model import StylometricClassifier
|
|
23
|
+
from ru_stylometry.vectorizer import StylometricVectorizer
|
|
24
|
+
|
|
25
|
+
__all__ = [
|
|
26
|
+
"ALL_GROUPS",
|
|
27
|
+
"DEFAULT_GROUPS",
|
|
28
|
+
"StylometricClassifier",
|
|
29
|
+
"StylometricVectorizer",
|
|
30
|
+
"describe_features",
|
|
31
|
+
"extract_features",
|
|
32
|
+
"feature_names",
|
|
33
|
+
]
|
|
34
|
+
__version__ = "0.1.1"
|
|
@@ -0,0 +1,76 @@
|
|
|
1
|
+
"""Small hand-made lexicons for lexicon-based features (lowercase, 'ё' folded into 'е')."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from typing import FrozenSet, Tuple
|
|
6
|
+
|
|
7
|
+
# Prepositions, conjunctions, particles, pronouns and forms of "быть". Stylometry has long used
|
|
8
|
+
# such function words as topic-independent markers of style (Mosteller & Wallace, 1964).
|
|
9
|
+
_FUNCTION_WORDS = """
|
|
10
|
+
в во на с со к ко по из изо от ото до за у о об обо для при про над надо под подо перед передо
|
|
11
|
+
между через без безо около после среди ради сквозь против кроме вместо внутри вне возле вокруг
|
|
12
|
+
мимо
|
|
13
|
+
и а но да или либо ни что чтобы чтоб как если когда пока потому поэтому так также тоже зато
|
|
14
|
+
однако хотя будто словно ведь
|
|
15
|
+
не бы же ли вот вон даже только лишь уже еще именно просто ну разве неужели почти пусть пускай
|
|
16
|
+
я ты он она оно мы вы они мой моя мое мои твой твоя твое твои его ее их наш наша наше наши
|
|
17
|
+
ваш ваша ваше ваши свой своя свое свои себя себе собой этот эта это эти тот та то те такой
|
|
18
|
+
такая такое такие весь вся все сам сама сами кто какой какая какое какие который которая
|
|
19
|
+
которое которые чей где куда откуда почему зачем чем чего чему кому кого ему ей им них ним нем
|
|
20
|
+
ней мне меня тебе тебя нас нам вас вам нами вами мной
|
|
21
|
+
быть был была было были буду будешь будет будем будете будут есть суть
|
|
22
|
+
"""
|
|
23
|
+
|
|
24
|
+
_NEGATIONS = """
|
|
25
|
+
не ни нет нельзя никогда никто ничто ничего никак нигде никуда никакой никакая никакое никакие
|
|
26
|
+
нисколько ничей
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
FUNCTION_WORDS: FrozenSet[str] = frozenset(_FUNCTION_WORDS.split())
|
|
30
|
+
NEGATIONS: FrozenSet[str] = frozenset(_NEGATIONS.split())
|
|
31
|
+
|
|
32
|
+
# Coordinating conjunctions that often open sentences in informal writing.
|
|
33
|
+
CONJUNCTION_STARTERS: FrozenSet[str] = frozenset("и а но да или либо зато ни".split())
|
|
34
|
+
|
|
35
|
+
# Connectives and stock phrases typical of formal prose and of machine-generated Russian text.
|
|
36
|
+
DISCOURSE_MARKERS: Tuple[str, ...] = (
|
|
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
|
+
)
|
ru_stylometry/_morph.py
ADDED
|
@@ -0,0 +1,41 @@
|
|
|
1
|
+
"""Cached morphological analysis on top of pymorphy3."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from functools import lru_cache
|
|
6
|
+
from typing import Any, NamedTuple, Optional
|
|
7
|
+
|
|
8
|
+
_analyzer: Any = None
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class WordInfo(NamedTuple):
|
|
12
|
+
"""The grammatical facts about one word that the morphological features need."""
|
|
13
|
+
|
|
14
|
+
pos: Optional[str]
|
|
15
|
+
case: Optional[str]
|
|
16
|
+
tense: Optional[str]
|
|
17
|
+
person: Optional[str]
|
|
18
|
+
lemma: str
|
|
19
|
+
is_known: bool
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def get_analyzer() -> Any:
|
|
23
|
+
"""Return a shared ``pymorphy3.MorphAnalyzer``; dictionaries are loaded on first use."""
|
|
24
|
+
global _analyzer
|
|
25
|
+
if _analyzer is None:
|
|
26
|
+
import pymorphy3
|
|
27
|
+
|
|
28
|
+
_analyzer = pymorphy3.MorphAnalyzer()
|
|
29
|
+
return _analyzer
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
@lru_cache(maxsize=100_000)
|
|
33
|
+
def analyze_word(word: str) -> WordInfo:
|
|
34
|
+
"""Analyse a lowercase word and keep its most probable parse.
|
|
35
|
+
|
|
36
|
+
Results are cached per process: word frequencies follow Zipf's law, so a few thousand
|
|
37
|
+
entries cover most of the running text and pymorphy3 parses each of them only once.
|
|
38
|
+
"""
|
|
39
|
+
parse = get_analyzer().parse(word)[0]
|
|
40
|
+
tag = parse.tag
|
|
41
|
+
return WordInfo(tag.POS, tag.case, tag.tense, tag.person, parse.normal_form, parse.is_known)
|
ru_stylometry/_stats.py
ADDED
|
@@ -0,0 +1,91 @@
|
|
|
1
|
+
"""Small numeric helpers shared by feature groups."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections import Counter
|
|
6
|
+
from typing import Hashable, Sequence
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def safe_div(numerator: float, denominator: float, default: float = 0.0) -> float:
|
|
10
|
+
"""Divide, returning ``default`` when the denominator is zero."""
|
|
11
|
+
return numerator / denominator if denominator else default
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def mean(values: Sequence[float]) -> float:
|
|
15
|
+
"""Arithmetic mean; 0.0 for an empty sequence."""
|
|
16
|
+
return sum(values) / len(values) if values else 0.0
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def pstdev(values: Sequence[float]) -> float:
|
|
20
|
+
"""Population standard deviation; 0.0 for an empty sequence."""
|
|
21
|
+
if not values:
|
|
22
|
+
return 0.0
|
|
23
|
+
mu = mean(values)
|
|
24
|
+
return (sum((v - mu) ** 2 for v in values) / len(values)) ** 0.5
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def mattr(tokens: Sequence[Hashable], window: int = 50) -> float:
|
|
28
|
+
"""Moving-average type-token ratio (Covington & McFall, 2010).
|
|
29
|
+
|
|
30
|
+
Mean share of distinct tokens over all windows of ``window`` consecutive tokens. Unlike the
|
|
31
|
+
plain type-token ratio it does not fall as the text gets longer. Texts shorter than the window
|
|
32
|
+
fall back to the plain type-token ratio.
|
|
33
|
+
"""
|
|
34
|
+
n = len(tokens)
|
|
35
|
+
if n == 0:
|
|
36
|
+
return 0.0
|
|
37
|
+
if n <= window:
|
|
38
|
+
return len(set(tokens)) / n
|
|
39
|
+
counts = Counter(tokens[:window])
|
|
40
|
+
unique = len(counts)
|
|
41
|
+
total = unique
|
|
42
|
+
for i in range(window, n):
|
|
43
|
+
leaving = tokens[i - window]
|
|
44
|
+
counts[leaving] -= 1
|
|
45
|
+
if counts[leaving] == 0:
|
|
46
|
+
unique -= 1
|
|
47
|
+
if counts[tokens[i]] == 0:
|
|
48
|
+
unique += 1
|
|
49
|
+
counts[tokens[i]] += 1
|
|
50
|
+
total += unique
|
|
51
|
+
return total / ((n - window + 1) * window)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _mtld_pass(tokens: Sequence[Hashable], threshold: float) -> float:
|
|
55
|
+
factors = 0.0
|
|
56
|
+
types: set = set()
|
|
57
|
+
count = 0
|
|
58
|
+
for token in tokens:
|
|
59
|
+
types.add(token)
|
|
60
|
+
count += 1
|
|
61
|
+
if len(types) / count <= threshold:
|
|
62
|
+
factors += 1.0
|
|
63
|
+
types = set()
|
|
64
|
+
count = 0
|
|
65
|
+
if count:
|
|
66
|
+
factors += (1.0 - len(types) / count) / (1.0 - threshold)
|
|
67
|
+
return len(tokens) / factors if factors else float(len(tokens))
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def mtld(tokens: Sequence[Hashable], threshold: float = 0.72) -> float:
|
|
71
|
+
"""Measure of textual lexical diversity (McCarthy & Jarvis, 2010).
|
|
72
|
+
|
|
73
|
+
Average length of a token run over which the type-token ratio stays above ``threshold``,
|
|
74
|
+
computed forwards and backwards. A text in which no run ever drops to the threshold gets a value
|
|
75
|
+
equal to its length.
|
|
76
|
+
"""
|
|
77
|
+
if not tokens:
|
|
78
|
+
return 0.0
|
|
79
|
+
forward = _mtld_pass(tokens, threshold)
|
|
80
|
+
backward = _mtld_pass(tokens[::-1], threshold)
|
|
81
|
+
return (forward + backward) / 2.0
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def yule_k(tokens: Sequence[Hashable]) -> float:
|
|
85
|
+
"""Yule's K characteristic (Yule, 1944): higher values mean more repetitive vocabulary."""
|
|
86
|
+
n = len(tokens)
|
|
87
|
+
if n == 0:
|
|
88
|
+
return 0.0
|
|
89
|
+
spectrum = Counter(Counter(tokens).values())
|
|
90
|
+
s2 = sum(m * m * v for m, v in spectrum.items())
|
|
91
|
+
return 1e4 * (s2 - n) / (n * n)
|
ru_stylometry/_text.py
ADDED
|
@@ -0,0 +1,99 @@
|
|
|
1
|
+
"""Low-level text helpers: tokenisation, sentence splitting and syllable counting."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import re
|
|
6
|
+
from typing import List
|
|
7
|
+
|
|
8
|
+
# A word is a run of letters, optionally joined by hyphens or apostrophes ("кто-то", "don't").
|
|
9
|
+
WORD_RE = re.compile(r"[^\W\d_]+(?:[-’'][^\W\d_]+)*")
|
|
10
|
+
CYRILLIC_RE = re.compile(r"[А-Яа-яЁё]")
|
|
11
|
+
|
|
12
|
+
# Words *longer* than this many letters are "long" (Björnsson's definition used by LIX).
|
|
13
|
+
LONG_WORD_LETTERS = 6
|
|
14
|
+
|
|
15
|
+
# Paragraph break, or a line break in front of a list item / heading marker.
|
|
16
|
+
_BLOCK_RE = re.compile(r"\n\s*\n|\n(?=[ \t]*(?:[-*•–—#]+|\d+[.)])[ \t]+)")
|
|
17
|
+
# Sentence terminator, optional closing quotes/brackets (captured) and the whitespace after them.
|
|
18
|
+
_BOUNDARY_RE = re.compile(r"(?<=[.!?…])([\"»”’)\]]*)\s+")
|
|
19
|
+
_INITIALS_RE = re.compile(r"^(?:[А-ЯЁA-Z]\.){1,3}$")
|
|
20
|
+
|
|
21
|
+
# Abbreviations after which a full stop does not end the sentence (lowercase, one token each).
|
|
22
|
+
_ABBREVIATION_LIST = """
|
|
23
|
+
т.д. т.п. т.е. т.к. т.н. т.о. т.ч. и.о. н.э. г. гг. в. вв. ул. им. пр. др. см. рис. табл. стр.
|
|
24
|
+
с. т. тт. тыс. млн. млрд. руб. коп. проф. доц. акад. канд. докт. напр. англ. лат. нем. франц.
|
|
25
|
+
ок. ср. изд. вып. гл. п. пп. ч. ст. обл. пос. д. кв. корп. оф. тел. тов.
|
|
26
|
+
"""
|
|
27
|
+
_ABBREVIATIONS = frozenset(_ABBREVIATION_LIST.split())
|
|
28
|
+
|
|
29
|
+
_VOWELS = frozenset("аеёиоуыэюяaeiouy")
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def word_length(word: str) -> int:
|
|
33
|
+
"""Number of letters in a word (hyphens and apostrophes are not counted)."""
|
|
34
|
+
return sum(1 for char in word if char.isalpha())
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def count_syllables(word: str) -> int:
|
|
38
|
+
"""Count syllables as vowel letters (Cyrillic and basic Latin vowels)."""
|
|
39
|
+
return sum(1 for char in word.lower() if char in _VOWELS)
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def truncate_words(text: str, max_words: int) -> str:
|
|
43
|
+
"""Keep the beginning of the text up to the end of its ``max_words``-th word.
|
|
44
|
+
|
|
45
|
+
Texts with at most ``max_words`` words are returned unchanged. Cutting before analysis bounds
|
|
46
|
+
the cost of long documents and puts texts of very different lengths on the same footing.
|
|
47
|
+
"""
|
|
48
|
+
end = 0
|
|
49
|
+
for count, match in enumerate(WORD_RE.finditer(text), start=1):
|
|
50
|
+
if count > max_words:
|
|
51
|
+
return text[:end]
|
|
52
|
+
end = match.end()
|
|
53
|
+
return text
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def _ends_with_abbreviation(chunk: str) -> bool:
|
|
57
|
+
last = chunk.rsplit(None, 1)[-1]
|
|
58
|
+
return last.lower() in _ABBREVIATIONS or bool(_INITIALS_RE.match(last))
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def _continues_after_ellipsis(previous: str, following: str) -> bool:
|
|
62
|
+
return previous.endswith(("…", "...")) and following[:1].islower()
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def _split_block(block: str) -> List[str]:
|
|
66
|
+
parts = _BOUNDARY_RE.split(block)
|
|
67
|
+
# parts = [sentence, closers, sentence, closers, ..., sentence]; glue closers back on.
|
|
68
|
+
chunks = [parts[i] + parts[i + 1] for i in range(0, len(parts) - 1, 2)] + [parts[-1]]
|
|
69
|
+
merged: List[str] = []
|
|
70
|
+
for chunk in chunks:
|
|
71
|
+
if merged and (
|
|
72
|
+
_ends_with_abbreviation(merged[-1]) or _continues_after_ellipsis(merged[-1], chunk)
|
|
73
|
+
):
|
|
74
|
+
merged[-1] = f"{merged[-1]} {chunk}"
|
|
75
|
+
else:
|
|
76
|
+
merged.append(chunk)
|
|
77
|
+
return merged
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
def split_sentences(text: str) -> List[str]:
|
|
81
|
+
"""Split Russian text into sentences with a rule-based splitter.
|
|
82
|
+
|
|
83
|
+
Boundaries are paragraph breaks, line breaks before list items, and sentence-final
|
|
84
|
+
punctuation followed by whitespace. A full stop after a known abbreviation ("т.д.", "г.") or
|
|
85
|
+
after initials ("А. С.") does not end a sentence, nor does an ellipsis followed by a lowercase
|
|
86
|
+
letter. Single line breaks inside a paragraph are treated as spaces.
|
|
87
|
+
|
|
88
|
+
Args:
|
|
89
|
+
text: Input text.
|
|
90
|
+
|
|
91
|
+
Returns:
|
|
92
|
+
Non-empty sentences in order of appearance.
|
|
93
|
+
"""
|
|
94
|
+
sentences: List[str] = []
|
|
95
|
+
for block in _BLOCK_RE.split(text):
|
|
96
|
+
block = " ".join(block.split())
|
|
97
|
+
if block:
|
|
98
|
+
sentences.extend(_split_block(block))
|
|
99
|
+
return sentences
|
ru_stylometry/cli.py
ADDED
|
@@ -0,0 +1,269 @@
|
|
|
1
|
+
"""Command-line interface of ``ru-stylometry``.
|
|
2
|
+
|
|
3
|
+
Commands: ``features`` and ``list-features`` work with the feature extractor; ``train``,
|
|
4
|
+
``evaluate``, ``predict`` and ``explain`` work with a classifier trained on a labelled table.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import argparse
|
|
10
|
+
import csv
|
|
11
|
+
import json
|
|
12
|
+
import sys
|
|
13
|
+
import time
|
|
14
|
+
from pathlib import Path
|
|
15
|
+
from typing import Any, Dict, List, Optional, Sequence, Tuple
|
|
16
|
+
|
|
17
|
+
from ru_stylometry import __version__
|
|
18
|
+
from ru_stylometry.evaluation import evaluate_classification, format_report
|
|
19
|
+
from ru_stylometry.features import (
|
|
20
|
+
ALL_GROUPS,
|
|
21
|
+
DEFAULT_GROUPS,
|
|
22
|
+
describe_features,
|
|
23
|
+
extract_features,
|
|
24
|
+
group_description,
|
|
25
|
+
)
|
|
26
|
+
from ru_stylometry.model import ESTIMATORS, StylometricClassifier
|
|
27
|
+
|
|
28
|
+
Row = Tuple[str, str, str]
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def _read_text(args: argparse.Namespace) -> str:
|
|
32
|
+
if args.file is not None:
|
|
33
|
+
return args.file.read_text(encoding="utf-8")
|
|
34
|
+
if args.text is not None:
|
|
35
|
+
return args.text
|
|
36
|
+
return sys.stdin.read()
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def _read_table(path: Path, text_column: str, label_column: str) -> Tuple[List[str], List[str]]:
|
|
40
|
+
"""Read texts and labels from a ``.csv``, ``.jsonl`` or ``.parquet`` file."""
|
|
41
|
+
suffix = path.suffix.lower()
|
|
42
|
+
if suffix == ".csv":
|
|
43
|
+
csv.field_size_limit(sys.maxsize)
|
|
44
|
+
with path.open(encoding="utf-8", newline="") as handle:
|
|
45
|
+
records: List[Dict[str, Any]] = list(csv.DictReader(handle))
|
|
46
|
+
elif suffix in (".jsonl", ".ndjson"):
|
|
47
|
+
lines = path.read_text(encoding="utf-8").splitlines()
|
|
48
|
+
records = [json.loads(line) for line in lines if line.strip()]
|
|
49
|
+
elif suffix == ".parquet":
|
|
50
|
+
try:
|
|
51
|
+
import pandas as pd
|
|
52
|
+
except ImportError as error:
|
|
53
|
+
raise ValueError(
|
|
54
|
+
"Reading Parquet needs pandas and pyarrow: pip install ru-stylometry[experiments]"
|
|
55
|
+
) from error
|
|
56
|
+
records = pd.read_parquet(path, columns=[text_column, label_column]).to_dict("records")
|
|
57
|
+
else:
|
|
58
|
+
raise ValueError(f"Unsupported file type {suffix!r}; use .csv, .jsonl or .parquet.")
|
|
59
|
+
for column in (text_column, label_column):
|
|
60
|
+
if not records or column not in records[0]:
|
|
61
|
+
available = sorted(records[0]) if records else []
|
|
62
|
+
raise ValueError(f"Column {column!r} not found in {path}; available: {available}.")
|
|
63
|
+
return [str(r[text_column]) for r in records], [str(r[label_column]) for r in records]
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def _markdown_catalogue(rows: Sequence[Row]) -> str:
|
|
67
|
+
by_group: Dict[str, List[Tuple[str, str]]] = {}
|
|
68
|
+
for group, name, description in rows:
|
|
69
|
+
by_group.setdefault(group, []).append((name, description))
|
|
70
|
+
default_count = sum(1 for group, _, _ in rows if group in DEFAULT_GROUPS)
|
|
71
|
+
lines = [
|
|
72
|
+
"# Feature catalogue",
|
|
73
|
+
"",
|
|
74
|
+
f"{len(rows)} features in {len(by_group)} groups. The {default_count} features of the "
|
|
75
|
+
"default groups form the default feature vector; groups marked *opt-in* are excluded "
|
|
76
|
+
"from it.",
|
|
77
|
+
]
|
|
78
|
+
for group, items in by_group.items():
|
|
79
|
+
status = "default" if group in DEFAULT_GROUPS else "opt-in"
|
|
80
|
+
lines += [
|
|
81
|
+
"",
|
|
82
|
+
f"## `{group}` ({status}, {len(items)} features)",
|
|
83
|
+
"",
|
|
84
|
+
group_description(group),
|
|
85
|
+
"",
|
|
86
|
+
"| Feature | Description |",
|
|
87
|
+
"|---|---|",
|
|
88
|
+
]
|
|
89
|
+
lines += [f"| `{name}` | {description} |" for name, description in items]
|
|
90
|
+
return "\n".join(lines)
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def _dump(data: Any) -> None:
|
|
94
|
+
print(json.dumps(data, ensure_ascii=False, indent=2))
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def _cmd_features(args: argparse.Namespace) -> int:
|
|
98
|
+
_dump(extract_features(_read_text(args), groups=args.groups, max_words=args.max_words))
|
|
99
|
+
return 0
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
def _cmd_list_features(args: argparse.Namespace) -> int:
|
|
103
|
+
rows = describe_features(args.groups or ALL_GROUPS)
|
|
104
|
+
if args.format == "json":
|
|
105
|
+
_dump([{"group": g, "name": n, "description": d} for g, n, d in rows])
|
|
106
|
+
elif args.format == "markdown":
|
|
107
|
+
print(_markdown_catalogue(rows))
|
|
108
|
+
else:
|
|
109
|
+
width = max(len(name) for _, name, _ in rows)
|
|
110
|
+
for group, name, description in rows:
|
|
111
|
+
print(f"{group:<12} {name:<{width}} {description}")
|
|
112
|
+
return 0
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
def _cmd_train(args: argparse.Namespace) -> int:
|
|
116
|
+
texts, labels = _read_table(args.data, args.text_column, args.label_column)
|
|
117
|
+
model = StylometricClassifier(
|
|
118
|
+
estimator=args.estimator,
|
|
119
|
+
groups=args.groups,
|
|
120
|
+
max_words=args.max_words,
|
|
121
|
+
class_weight="balanced" if args.balanced else None,
|
|
122
|
+
n_jobs=args.n_jobs,
|
|
123
|
+
)
|
|
124
|
+
started = time.perf_counter()
|
|
125
|
+
model.fit(texts, labels)
|
|
126
|
+
summary: Dict[str, Any] = {
|
|
127
|
+
"model": str(args.out),
|
|
128
|
+
"classes": [str(c) for c in model.classes_],
|
|
129
|
+
"texts": len(texts),
|
|
130
|
+
"features": len(model.feature_names_),
|
|
131
|
+
"seconds": round(time.perf_counter() - started, 1),
|
|
132
|
+
}
|
|
133
|
+
model.save(str(args.out))
|
|
134
|
+
if args.valid is not None:
|
|
135
|
+
valid_texts, valid_labels = _read_table(args.valid, args.text_column, args.label_column)
|
|
136
|
+
report = evaluate_classification(valid_labels, model.predict(valid_texts), model.classes_)
|
|
137
|
+
summary["valid"] = {k: report[k] for k in ("n", "accuracy", "macro_f1")}
|
|
138
|
+
_dump(summary)
|
|
139
|
+
return 0
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
def _cmd_evaluate(args: argparse.Namespace) -> int:
|
|
143
|
+
model = StylometricClassifier.load(str(args.model))
|
|
144
|
+
texts, labels = _read_table(args.data, args.text_column, args.label_column)
|
|
145
|
+
started = time.perf_counter()
|
|
146
|
+
predicted = model.predict(texts)
|
|
147
|
+
seconds = time.perf_counter() - started
|
|
148
|
+
report = evaluate_classification(labels, predicted, labels=[str(c) for c in model.classes_])
|
|
149
|
+
report["ms_per_document"] = 1000.0 * seconds / max(len(texts), 1)
|
|
150
|
+
if args.format == "json":
|
|
151
|
+
_dump(report)
|
|
152
|
+
else:
|
|
153
|
+
print(format_report(report))
|
|
154
|
+
print(f"\nTime per document: {report['ms_per_document']:.1f} ms")
|
|
155
|
+
return 0
|
|
156
|
+
|
|
157
|
+
|
|
158
|
+
def _cmd_predict(args: argparse.Namespace) -> int:
|
|
159
|
+
model = StylometricClassifier.load(str(args.model))
|
|
160
|
+
text = _read_text(args)
|
|
161
|
+
probabilities = model.predict_proba([text])[0]
|
|
162
|
+
best = int(probabilities.argmax())
|
|
163
|
+
_dump(
|
|
164
|
+
{
|
|
165
|
+
"label": str(model.classes_[best]),
|
|
166
|
+
"probability": float(probabilities[best]),
|
|
167
|
+
"probabilities": {str(c): float(p) for c, p in zip(model.classes_, probabilities)},
|
|
168
|
+
}
|
|
169
|
+
)
|
|
170
|
+
return 0
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
def _cmd_explain(args: argparse.Namespace) -> int:
|
|
174
|
+
model = StylometricClassifier.load(str(args.model))
|
|
175
|
+
explanation = model.explain(_read_text(args), target=args.target, top_k=args.top)
|
|
176
|
+
if args.format == "json":
|
|
177
|
+
_dump(explanation)
|
|
178
|
+
return 0
|
|
179
|
+
print(f"label: {explanation['label']} (probability {explanation['probability']:.3f})")
|
|
180
|
+
print(f"effect on the log-odds of class '{explanation['target']}' (positive: towards it):")
|
|
181
|
+
for item in explanation["features"]:
|
|
182
|
+
print(
|
|
183
|
+
f" {item['feature']:<26} {item['effect']:+.2f} value {item['value']:.3f}, "
|
|
184
|
+
f"typical {item['typical']:.3f} [{item['group']}]"
|
|
185
|
+
)
|
|
186
|
+
print(
|
|
187
|
+
"by group: "
|
|
188
|
+
+ ", ".join(f"{g} {v['effect']:+.2f}" for g, v in explanation["groups"].items())
|
|
189
|
+
)
|
|
190
|
+
return 0
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
def _add_text_input(parser: argparse.ArgumentParser) -> None:
|
|
194
|
+
parser.add_argument("text", nargs="?", help="text to analyse (default: read from stdin)")
|
|
195
|
+
parser.add_argument("-f", "--file", type=Path, help="read the text from a UTF-8 file")
|
|
196
|
+
|
|
197
|
+
|
|
198
|
+
def _add_table_columns(parser: argparse.ArgumentParser) -> None:
|
|
199
|
+
parser.add_argument("--text-column", default="text", help="column with the texts")
|
|
200
|
+
parser.add_argument("--label-column", default="label", help="column with the class labels")
|
|
201
|
+
|
|
202
|
+
|
|
203
|
+
def _build_parser() -> argparse.ArgumentParser:
|
|
204
|
+
parser = argparse.ArgumentParser(
|
|
205
|
+
prog="ru-stylometry",
|
|
206
|
+
description="Interpretable stylometric features and classifiers for Russian text.",
|
|
207
|
+
)
|
|
208
|
+
parser.add_argument("--version", action="version", version=f"%(prog)s {__version__}")
|
|
209
|
+
commands = parser.add_subparsers(dest="command", required=True)
|
|
210
|
+
|
|
211
|
+
features = commands.add_parser("features", help="extract features from a text as JSON")
|
|
212
|
+
_add_text_input(features)
|
|
213
|
+
features.add_argument("-g", "--groups", nargs="+", metavar="GROUP", help="feature groups")
|
|
214
|
+
features.add_argument(
|
|
215
|
+
"--max-words", type=int, metavar="N", help="analyse only the first N words"
|
|
216
|
+
)
|
|
217
|
+
features.set_defaults(handler=_cmd_features)
|
|
218
|
+
|
|
219
|
+
catalogue = commands.add_parser("list-features", help="print the feature catalogue")
|
|
220
|
+
catalogue.add_argument("-g", "--groups", nargs="+", metavar="GROUP", help="feature groups")
|
|
221
|
+
catalogue.add_argument("--format", choices=["text", "markdown", "json"], default="text")
|
|
222
|
+
catalogue.set_defaults(handler=_cmd_list_features)
|
|
223
|
+
|
|
224
|
+
train = commands.add_parser("train", help="train a classifier on a labelled table")
|
|
225
|
+
train.add_argument(
|
|
226
|
+
"data", type=Path, help=".csv, .jsonl or .parquet file with texts and labels"
|
|
227
|
+
)
|
|
228
|
+
train.add_argument("-o", "--out", type=Path, required=True, help="where to save the model")
|
|
229
|
+
_add_table_columns(train)
|
|
230
|
+
train.add_argument("--estimator", choices=ESTIMATORS, default="hgb")
|
|
231
|
+
train.add_argument("-g", "--groups", nargs="+", metavar="GROUP", help="feature groups")
|
|
232
|
+
train.add_argument("--max-words", type=int, default=500, metavar="N")
|
|
233
|
+
train.add_argument("--balanced", action="store_true", help="weight classes equally")
|
|
234
|
+
train.add_argument("--n-jobs", type=int, help="worker processes for feature extraction")
|
|
235
|
+
train.add_argument("--valid", type=Path, help="labelled file to score the model on")
|
|
236
|
+
train.set_defaults(handler=_cmd_train)
|
|
237
|
+
|
|
238
|
+
evaluate = commands.add_parser("evaluate", help="score a saved model on a labelled table")
|
|
239
|
+
evaluate.add_argument("data", type=Path, help=".csv, .jsonl or .parquet file")
|
|
240
|
+
evaluate.add_argument("-m", "--model", type=Path, required=True)
|
|
241
|
+
_add_table_columns(evaluate)
|
|
242
|
+
evaluate.add_argument("--format", choices=["markdown", "json"], default="markdown")
|
|
243
|
+
evaluate.set_defaults(handler=_cmd_evaluate)
|
|
244
|
+
|
|
245
|
+
predict = commands.add_parser("predict", help="classify one text")
|
|
246
|
+
_add_text_input(predict)
|
|
247
|
+
predict.add_argument("-m", "--model", type=Path, required=True)
|
|
248
|
+
predict.set_defaults(handler=_cmd_predict)
|
|
249
|
+
|
|
250
|
+
explain = commands.add_parser("explain", help="explain the decision for one text")
|
|
251
|
+
_add_text_input(explain)
|
|
252
|
+
explain.add_argument("-m", "--model", type=Path, required=True)
|
|
253
|
+
explain.add_argument("--target", help="class to explain (default: the predicted class)")
|
|
254
|
+
explain.add_argument("--top", type=int, default=10, help="number of features to list")
|
|
255
|
+
explain.add_argument("--format", choices=["text", "json"], default="text")
|
|
256
|
+
explain.set_defaults(handler=_cmd_explain)
|
|
257
|
+
|
|
258
|
+
parser.epilog = f"feature groups: {', '.join(ALL_GROUPS)}"
|
|
259
|
+
return parser
|
|
260
|
+
|
|
261
|
+
|
|
262
|
+
def main(argv: Optional[Sequence[str]] = None) -> int:
|
|
263
|
+
"""Entry point of the ``ru-stylometry`` command."""
|
|
264
|
+
args = _build_parser().parse_args(argv)
|
|
265
|
+
try:
|
|
266
|
+
return args.handler(args)
|
|
267
|
+
except (ValueError, TypeError, OSError) as error:
|
|
268
|
+
print(f"ru-stylometry: error: {error}", file=sys.stderr)
|
|
269
|
+
return 2
|