pytextad 0.1.0__tar.gz

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.
pytextad-0.1.0/LICENSE ADDED
@@ -0,0 +1,24 @@
1
+ BSD 2-Clause License
2
+
3
+ Copyright (c) 2026, Yang Cao
4
+
5
+ Redistribution and use in source and binary forms, with or without
6
+ modification, are permitted provided that the following conditions are met:
7
+
8
+ 1. Redistributions of source code must retain the above copyright notice, this
9
+ list of conditions and the following disclaimer.
10
+
11
+ 2. Redistributions in binary form must reproduce the above copyright notice,
12
+ this list of conditions and the following disclaimer in the documentation
13
+ and/or other materials provided with the distribution.
14
+
15
+ THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
16
+ AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
17
+ IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
18
+ DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
19
+ FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
20
+ DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
21
+ SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
22
+ CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
23
+ OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
24
+ OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
@@ -0,0 +1,108 @@
1
+ Metadata-Version: 2.4
2
+ Name: pytextad
3
+ Version: 0.1.0
4
+ Summary: A unified library for text anomaly detection (document- and token-level), in the style of PyOD
5
+ Author: Yang Cao
6
+ License-Expression: BSD-2-Clause
7
+ Project-URL: Homepage, https://github.com/charles-cao/pytextad
8
+ Project-URL: Documentation, https://pytextad.readthedocs.io
9
+ Keywords: anomaly detection,outlier detection,text,NLP,token-level
10
+ Classifier: Development Status :: 3 - Alpha
11
+ Classifier: Intended Audience :: Science/Research
12
+ Classifier: Programming Language :: Python :: 3
13
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
14
+ Requires-Python: >=3.9
15
+ Description-Content-Type: text/markdown
16
+ License-File: LICENSE
17
+ License-File: THIRD_PARTY_NOTICES.md
18
+ Requires-Dist: numpy>=1.21
19
+ Requires-Dist: scikit-learn>=1.0
20
+ Requires-Dist: torch>=1.13
21
+ Requires-Dist: transformers>=4.30
22
+ Provides-Extra: test
23
+ Requires-Dist: pytest>=7; extra == "test"
24
+ Provides-Extra: docs
25
+ Requires-Dist: sphinx>=7; extra == "docs"
26
+ Requires-Dist: furo; extra == "docs"
27
+ Requires-Dist: myst-parser; extra == "docs"
28
+ Dynamic: license-file
29
+
30
+ # PyTextAD: Text Anomaly Detection in Python
31
+
32
+ [![PyPI](https://img.shields.io/pypi/v/pytextad.svg)](https://pypi.org/project/pytextad/)
33
+ [![Documentation](https://readthedocs.org/projects/pytextad/badge/?version=latest)](https://pytextad.readthedocs.io)
34
+ [![Tests](https://github.com/charles-cao/pytextad/actions/workflows/tests.yml/badge.svg)](https://github.com/charles-cao/pytextad/actions/workflows/tests.yml)
35
+ [![License](https://img.shields.io/badge/license-BSD--2--Clause-blue.svg)](LICENSE)
36
+
37
+ **PyTextAD** is a Python library for detecting anomalies in text, at the document
38
+ and at the token level. Every detector follows the PyOD interface
39
+ (`fit`, `decision_function`, `predict`, `decision_scores_`, `labels_`), and every
40
+ re-implemented method is checked numerically against its original code.
41
+
42
+ ## Installation
43
+
44
+ ```bash
45
+ pip install pytextad
46
+ ```
47
+
48
+ From source:
49
+
50
+ ```bash
51
+ git clone https://github.com/charles-cao/pytextad.git
52
+ cd pytextad
53
+ pip install -e ".[test]"
54
+ ```
55
+
56
+ Requires Python >= 3.9, PyTorch >= 1.13 and transformers >= 4.30. Install the
57
+ PyTorch build that matches your CUDA version first (https://pytorch.org).
58
+
59
+ ## Quick start
60
+
61
+ ```python
62
+ from pytextad import CVDD, DATE, FATE, RSRAE, TokenEmbedder, mean_pool
63
+
64
+ # frozen token embeddings from any Hugging Face encoder
65
+ emb = TokenEmbedder("bert-base-uncased")
66
+ H_train, _ = emb.transform(train_texts)
67
+ H_test, _ = emb.transform(test_texts)
68
+
69
+ scores = CVDD().fit(H_train).decision_function(H_test) # higher = more anomalous
70
+ scores = RSRAE().fit(mean_pool(H_train)).decision_function(mean_pool(H_test))
71
+ scores = DATE().fit(train_texts).decision_function(test_texts) # raw text in, trains its own model
72
+ scores = FATE().fit(texts, y).decision_function(test_texts) # few-shot: y = 1 for labelled anomalies
73
+
74
+ word_scores = DATE().fit(train_texts).token_scores([t.split() for t in test_texts])
75
+ ```
76
+
77
+ A runnable example on AG News: `python examples/quickstart.py`.
78
+
79
+ ## Implemented methods
80
+
81
+ | Method | Year | Input | Token scores | Reference |
82
+ |---|---|---|---|---|
83
+ | CVDD | 2019 | frozen token embeddings | yes | Ruff et al., *Self-Attentive, Multi-Context One-Class Classification for Unsupervised Anomaly Detection on Text*, ACL 2019 |
84
+ | RSRAE | 2020 | document vectors | no | Lai et al., *Robust Subspace Recovery Layer for Unsupervised Anomaly Detection*, ICLR 2020 |
85
+ | DATE | 2021 | raw text | yes | Manolache et al., *DATE: Detecting Anomalies in Text via Self-Supervision of Transformers*, NAACL 2021 |
86
+ | FATE | 2023 | raw text (+ few labelled anomalies) | no | Das et al., *Few-shot Anomaly Detection in Text with Deviation Learning*, ICONIP 2023 |
87
+
88
+ Default hyperparameters are those of the official code. Each module's docstring
89
+ lists where the official code and the paper disagree and which one we follow.
90
+
91
+ ## Faithfulness to the original implementations
92
+
93
+ `tests/verification/` runs each original implementation next to ours with the same
94
+ weights, inputs and random seeds and compares the results (DATE against the original
95
+ transformers 3.0.2 code, RSRAE against the original TensorFlow code). All checks
96
+ pass; see [tests/verification/README.md](tests/verification/README.md).
97
+
98
+ ## Running the tests
99
+
100
+ ```bash
101
+ pytest # fast API tests, a few seconds
102
+ PYTEXTAD_DEVICE=cuda pytest # same, on GPU (PowerShell: $env:PYTEXTAD_DEVICE="cuda"; pytest)
103
+ ```
104
+
105
+ ## License
106
+
107
+ BSD 2-Clause. Third-party notices for the original implementations are in
108
+ [THIRD_PARTY_NOTICES.md](THIRD_PARTY_NOTICES.md).
@@ -0,0 +1,79 @@
1
+ # PyTextAD: Text Anomaly Detection in Python
2
+
3
+ [![PyPI](https://img.shields.io/pypi/v/pytextad.svg)](https://pypi.org/project/pytextad/)
4
+ [![Documentation](https://readthedocs.org/projects/pytextad/badge/?version=latest)](https://pytextad.readthedocs.io)
5
+ [![Tests](https://github.com/charles-cao/pytextad/actions/workflows/tests.yml/badge.svg)](https://github.com/charles-cao/pytextad/actions/workflows/tests.yml)
6
+ [![License](https://img.shields.io/badge/license-BSD--2--Clause-blue.svg)](LICENSE)
7
+
8
+ **PyTextAD** is a Python library for detecting anomalies in text, at the document
9
+ and at the token level. Every detector follows the PyOD interface
10
+ (`fit`, `decision_function`, `predict`, `decision_scores_`, `labels_`), and every
11
+ re-implemented method is checked numerically against its original code.
12
+
13
+ ## Installation
14
+
15
+ ```bash
16
+ pip install pytextad
17
+ ```
18
+
19
+ From source:
20
+
21
+ ```bash
22
+ git clone https://github.com/charles-cao/pytextad.git
23
+ cd pytextad
24
+ pip install -e ".[test]"
25
+ ```
26
+
27
+ Requires Python >= 3.9, PyTorch >= 1.13 and transformers >= 4.30. Install the
28
+ PyTorch build that matches your CUDA version first (https://pytorch.org).
29
+
30
+ ## Quick start
31
+
32
+ ```python
33
+ from pytextad import CVDD, DATE, FATE, RSRAE, TokenEmbedder, mean_pool
34
+
35
+ # frozen token embeddings from any Hugging Face encoder
36
+ emb = TokenEmbedder("bert-base-uncased")
37
+ H_train, _ = emb.transform(train_texts)
38
+ H_test, _ = emb.transform(test_texts)
39
+
40
+ scores = CVDD().fit(H_train).decision_function(H_test) # higher = more anomalous
41
+ scores = RSRAE().fit(mean_pool(H_train)).decision_function(mean_pool(H_test))
42
+ scores = DATE().fit(train_texts).decision_function(test_texts) # raw text in, trains its own model
43
+ scores = FATE().fit(texts, y).decision_function(test_texts) # few-shot: y = 1 for labelled anomalies
44
+
45
+ word_scores = DATE().fit(train_texts).token_scores([t.split() for t in test_texts])
46
+ ```
47
+
48
+ A runnable example on AG News: `python examples/quickstart.py`.
49
+
50
+ ## Implemented methods
51
+
52
+ | Method | Year | Input | Token scores | Reference |
53
+ |---|---|---|---|---|
54
+ | CVDD | 2019 | frozen token embeddings | yes | Ruff et al., *Self-Attentive, Multi-Context One-Class Classification for Unsupervised Anomaly Detection on Text*, ACL 2019 |
55
+ | RSRAE | 2020 | document vectors | no | Lai et al., *Robust Subspace Recovery Layer for Unsupervised Anomaly Detection*, ICLR 2020 |
56
+ | DATE | 2021 | raw text | yes | Manolache et al., *DATE: Detecting Anomalies in Text via Self-Supervision of Transformers*, NAACL 2021 |
57
+ | FATE | 2023 | raw text (+ few labelled anomalies) | no | Das et al., *Few-shot Anomaly Detection in Text with Deviation Learning*, ICONIP 2023 |
58
+
59
+ Default hyperparameters are those of the official code. Each module's docstring
60
+ lists where the official code and the paper disagree and which one we follow.
61
+
62
+ ## Faithfulness to the original implementations
63
+
64
+ `tests/verification/` runs each original implementation next to ours with the same
65
+ weights, inputs and random seeds and compares the results (DATE against the original
66
+ transformers 3.0.2 code, RSRAE against the original TensorFlow code). All checks
67
+ pass; see [tests/verification/README.md](tests/verification/README.md).
68
+
69
+ ## Running the tests
70
+
71
+ ```bash
72
+ pytest # fast API tests, a few seconds
73
+ PYTEXTAD_DEVICE=cuda pytest # same, on GPU (PowerShell: $env:PYTEXTAD_DEVICE="cuda"; pytest)
74
+ ```
75
+
76
+ ## License
77
+
78
+ BSD 2-Clause. Third-party notices for the original implementations are in
79
+ [THIRD_PARTY_NOTICES.md](THIRD_PARTY_NOTICES.md).
@@ -0,0 +1,36 @@
1
+ # Third-party notices
2
+
3
+ ## CVDD (pytextad/cvdd.py)
4
+ Re-implemented from https://github.com/lukasruff/CVDD-PyTorch
5
+
6
+ MIT License
7
+
8
+ Copyright (c) 2019 lukasruff
9
+
10
+ Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the "Software"), to deal in the Software without restriction, including without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all copies or substantial portions of the Software.
13
+
14
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
15
+
16
+ ## RSRAE (pytextad/rsrae.py)
17
+ Ported to PyTorch from https://github.com/dmzou/RSRAE
18
+
19
+ MIT License
20
+
21
+ Copyright (c) 2019-present Chieh-Hsin Lai, Dongmian Zou and Gilad Lerman
22
+
23
+ Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the "Software"), to deal in the Software without restriction, including without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions:
24
+
25
+ The above copyright notice and this permission notice shall be included in all copies or substantial portions of the Software.
26
+
27
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
28
+
29
+ ## DATE (pytextad/date.py)
30
+ Re-implemented from https://github.com/bit-ml/date, whose model code lives in a
31
+ modified copy of simpletransformers distributed under the Apache License 2.0.
32
+ No source code is copied; the logic was re-written against modern transformers.
33
+
34
+ ## FATE (pytextad/fate.py)
35
+ Re-implemented from https://github.com/arav1ndajay/fate, which has no licence file.
36
+ No source code is copied; the logic was re-written and checked numerically against it.
@@ -0,0 +1,44 @@
1
+ [build-system]
2
+ requires = ["setuptools>=77"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [project]
6
+ name = "pytextad"
7
+ dynamic = ["version"]
8
+ description = "A unified library for text anomaly detection (document- and token-level), in the style of PyOD"
9
+ readme = "README.md"
10
+ license = "BSD-2-Clause"
11
+ license-files = ["LICENSE", "THIRD_PARTY_NOTICES.md"]
12
+ authors = [{name = "Yang Cao"}]
13
+ requires-python = ">=3.9"
14
+ keywords = ["anomaly detection", "outlier detection", "text", "NLP", "token-level"]
15
+ classifiers = [
16
+ "Development Status :: 3 - Alpha",
17
+ "Intended Audience :: Science/Research",
18
+ "Programming Language :: Python :: 3",
19
+ "Topic :: Scientific/Engineering :: Artificial Intelligence",
20
+ ]
21
+ dependencies = [
22
+ "numpy>=1.21",
23
+ "scikit-learn>=1.0",
24
+ "torch>=1.13",
25
+ "transformers>=4.30",
26
+ ]
27
+
28
+ [project.optional-dependencies]
29
+ test = ["pytest>=7"]
30
+ docs = ["sphinx>=7", "furo", "myst-parser"]
31
+
32
+ [project.urls]
33
+ Homepage = "https://github.com/charles-cao/pytextad"
34
+ Documentation = "https://pytextad.readthedocs.io"
35
+
36
+ [tool.setuptools.dynamic]
37
+ version = {attr = "pytextad.version.__version__"}
38
+
39
+ [tool.setuptools.packages.find]
40
+ include = ["pytextad*"]
41
+
42
+ [tool.pytest.ini_options]
43
+ testpaths = ["tests"]
44
+ addopts = "--ignore=tests/verification"
@@ -0,0 +1,9 @@
1
+ """PyTextAD: a unified library for text anomaly detection, in the style of PyOD."""
2
+ from .version import __version__
3
+ from .models.cvdd import CVDD
4
+ from .models.date import DATE
5
+ from .models.fate import FATE
6
+ from .models.rsrae import RSRAE
7
+ from .utils.embeddings import TokenEmbedder, mean_pool, words_from_subwords
8
+
9
+ __all__ = ["__version__", "CVDD", "DATE", "FATE", "RSRAE", "TokenEmbedder", "mean_pool", "words_from_subwords"]
@@ -0,0 +1,7 @@
1
+ from .base import BaseTextDetector
2
+ from .cvdd import CVDD
3
+ from .date import DATE
4
+ from .fate import FATE
5
+ from .rsrae import RSRAE
6
+
7
+ __all__ = ["BaseTextDetector", "CVDD", "DATE", "FATE", "RSRAE"]
@@ -0,0 +1,71 @@
1
+ """Common interface for all detectors, modelled on PyOD's BaseDetector.
2
+
3
+ Conventions (identical to PyOD):
4
+ * fit(X, y=None) returns self and sets ``decision_scores_`` (scores of the training data),
5
+ ``threshold_`` and ``labels_``.
6
+ * decision_function(X) returns one score per sample; higher = more anomalous.
7
+ * predict(X) returns 0/1 using ``threshold_`` (the (1 - contamination) quantile of
8
+ the training scores).
9
+ Detectors that can score tokens additionally implement ``token_scores``.
10
+ """
11
+
12
+ import abc
13
+ import random
14
+
15
+ import numpy as np
16
+ import torch
17
+
18
+
19
+ class BaseTextDetector(abc.ABC):
20
+
21
+ def __init__(self, contamination=0.1, random_state=0, device=None, verbose=False):
22
+ if not 0.0 < contamination <= 0.5:
23
+ raise ValueError("contamination must be in (0, 0.5]")
24
+ self.contamination = contamination
25
+ self.random_state = random_state
26
+ self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
27
+ self.verbose = verbose
28
+
29
+ # ------------------------------------------------------------------ to implement
30
+ @abc.abstractmethod
31
+ def fit(self, X, y=None):
32
+ """Fit the detector on (mostly) normal data and set ``decision_scores_``,
33
+ ``threshold_`` and ``labels_``. ``y`` is ignored except by semi-supervised
34
+ detectors (FATE). Returns ``self``."""
35
+
36
+ @abc.abstractmethod
37
+ def decision_function(self, X):
38
+ """Anomaly score of every sample in ``X``; higher means more anomalous."""
39
+
40
+ # ------------------------------------------------------------------ shared
41
+ def _set_seed(self):
42
+ if self.random_state is None:
43
+ return
44
+ random.seed(self.random_state)
45
+ np.random.seed(self.random_state)
46
+ torch.manual_seed(self.random_state)
47
+ if torch.cuda.is_available():
48
+ torch.cuda.manual_seed_all(self.random_state)
49
+
50
+ def _process_decision_scores(self, scores):
51
+ self.decision_scores_ = np.asarray(scores, dtype=float)
52
+ self.threshold_ = float(np.percentile(self.decision_scores_, 100 * (1 - self.contamination)))
53
+ self.labels_ = (self.decision_scores_ > self.threshold_).astype(int)
54
+ return self
55
+
56
+ def predict(self, X):
57
+ """Binary labels (1 = anomaly) using ``threshold_`` from the training scores."""
58
+ self._check_fitted()
59
+ return (self.decision_function(X) > self.threshold_).astype(int)
60
+
61
+ def fit_predict(self, X, y=None):
62
+ """Fit on ``X`` and return the labels of the training data (``labels_``)."""
63
+ return self.fit(X, y).labels_
64
+
65
+ def _check_fitted(self):
66
+ if not hasattr(self, "decision_scores_"):
67
+ raise RuntimeError(f"{type(self).__name__} is not fitted yet; call fit() first.")
68
+
69
+ def _log(self, msg):
70
+ if self.verbose:
71
+ print(f"[{type(self).__name__}] {msg}")
@@ -0,0 +1,221 @@
1
+ """CVDD: Context Vector Data Description (Ruff et al., ACL 2019).
2
+
3
+ Re-implementation for modern PyTorch. The model and training logic follow the
4
+ official code, https://github.com/lukasruff/CVDD-PyTorch (MIT License,
5
+ Copyright (c) 2019 lukasruff); see THIRD_PARTY_NOTICES.
6
+
7
+ Input: a list of frozen token-embedding matrices, one [n_tokens, dim] array per
8
+ document (GloVe/fastText vectors, or hidden states of a frozen PLM, e.g. from
9
+ ``pytextad.embeddings.TokenEmbedder``). Only the self-attention layer and the context
10
+ vectors are trained, exactly as in the official code.
11
+
12
+ Faithful to the official code
13
+ * self-attention A = softmax_over_tokens(W2 tanh(W1 H)), r heads, no biases
14
+ * M = A H; cosine distance d_k = 0.5 (1 - cos(M_k, c_k))
15
+ * context vectors initialised by k-means on L2-normalised mean token embeddings
16
+ of the training set, then L2-normalised
17
+ * loss = mean_n sum_k softmax_k(-alpha d) d_k + lambda_p * mean((C C^T - I)^2)
18
+ * temperature alpha annealed at 5 equidistant milestones
19
+ (soft / linear / logarithmic / hard, same values as the official code)
20
+ * Adam with weight decay, gradient-norm clipping at 0.5, MultiStepLR(gamma=0.1)
21
+ stepped at the start of each epoch as in the official trainer (so with
22
+ PyTorch >= 1.1 the LR drop happens one epoch before the nominal milestone)
23
+ * anomaly score = mean_k d_k ("context_dist_mean", the official default)
24
+ * defaults = the settings in the official README for Reuters / 20 Newsgroups
25
+ (3 heads, attention size 150, lambda_p 1, logarithmic, 100 epochs, lr 0.01,
26
+ lr milestone 40, batch 64, weight decay 5e-7)
27
+
28
+ Deliberate deviations
29
+ * Padding is masked out of the attention softmax. The official code pads with index
30
+ 0 and does not mask. With static word vectors (GloVe, fastText) the pad vector is
31
+ zero, which only rescales M, so the cosine distances, the loss and the training
32
+ are identical (verified numerically with variable-length documents). With PLM
33
+ hidden states the official "bert" option feeds non-zero pad vectors into M, which
34
+ makes a document's score depend on the length of the others in its batch; we do
35
+ not reproduce that.
36
+ * Mini-batches are drawn uniformly at random and every document is scored. The
37
+ official loaders group documents by length (BucketBatchSampler) with
38
+ drop_last=True for both training and testing, so up to batch_size - 1 test
39
+ documents are never scored in the official evaluation.
40
+ * The official "context_best" score is NOT provided: it selects, per run, the
41
+ head with the highest AUROC on the labelled test set, which uses test labels.
42
+ ``head_scores()`` returns per-head distances if you need them.
43
+
44
+ Extension (not in the CVDD paper)
45
+ * token_scores(): min over heads of the cosine distance between each token vector
46
+ and the context vectors. Comparable across documents. Use with care: it is our
47
+ definition, not part of CVDD.
48
+ """
49
+
50
+ import warnings
51
+
52
+ import numpy as np
53
+ import torch
54
+ import torch.nn as nn
55
+ import torch.nn.functional as F
56
+ from sklearn.cluster import KMeans
57
+
58
+ from .base import BaseTextDetector
59
+
60
+
61
+ def _alpha_schedule(name, n_epochs):
62
+ milestones = list(np.arange(1, 6) * int(n_epochs / 5))
63
+ if name == "soft":
64
+ alphas = [0.0] * 5
65
+ elif name == "linear":
66
+ alphas = list(np.linspace(0.2, 1, 5))
67
+ elif name == "logarithmic":
68
+ alphas = list(np.logspace(-4, 0, 5))
69
+ elif name == "hard":
70
+ alphas = [100.0] * 4 # official list has 4 entries
71
+ else:
72
+ raise ValueError(f"unknown alpha_scheduler {name}")
73
+ return milestones, alphas
74
+
75
+
76
+ class _CVDDNet(nn.Module):
77
+ def __init__(self, dim, attention_size, n_heads):
78
+ super().__init__()
79
+ self.W1 = nn.Linear(dim, attention_size, bias=False)
80
+ self.W2 = nn.Linear(attention_size, n_heads, bias=False)
81
+ self.c = nn.Parameter((torch.rand(n_heads, dim) - 0.5) * 2)
82
+ self.alpha = 0.0
83
+
84
+ def attend(self, H, mask):
85
+ # H [B, L, d], mask [B, L] (True = real token)
86
+ logits = self.W2(torch.tanh(self.W1(H))) # [B, L, r]
87
+ logits = logits.masked_fill(~mask.unsqueeze(-1), float("-inf"))
88
+ A = F.softmax(logits, dim=1).transpose(1, 2) # [B, r, L]
89
+ return A @ H, A # M [B, r, d]
90
+
91
+ def forward(self, H, mask):
92
+ M, A = self.attend(H, mask)
93
+ d = 0.5 * (1 - F.cosine_similarity(M, self.c.unsqueeze(0), dim=2)) # [B, r]
94
+ w = F.softmax(-self.alpha * d, dim=1)
95
+ return d, w, A
96
+
97
+
98
+ class CVDD(BaseTextDetector):
99
+
100
+ def __init__(self, n_heads=3, attention_size=150, lambda_p=1.0,
101
+ alpha_scheduler="logarithmic", n_epochs=100, lr=0.01, lr_milestones=(40,),
102
+ batch_size=64, weight_decay=0.5e-6, contamination=0.1, random_state=0,
103
+ device=None, verbose=False):
104
+ super().__init__(contamination, random_state, device, verbose)
105
+ self.n_heads = n_heads
106
+ self.attention_size = attention_size
107
+ self.lambda_p = lambda_p
108
+ self.alpha_scheduler = alpha_scheduler
109
+ self.n_epochs = n_epochs
110
+ self.lr = lr
111
+ self.lr_milestones = tuple(lr_milestones)
112
+ self.batch_size = batch_size
113
+ self.weight_decay = weight_decay
114
+
115
+ # ------------------------------------------------------------------ batching
116
+ @staticmethod
117
+ def _check(X):
118
+ X = [np.asarray(x, dtype=np.float32) for x in X]
119
+ if any(x.ndim != 2 or len(x) == 0 for x in X):
120
+ raise ValueError("X must be a list of non-empty [n_tokens, dim] arrays")
121
+ if len({x.shape[1] for x in X}) != 1:
122
+ raise ValueError("all token embeddings must have the same dimension")
123
+ return X
124
+
125
+ def _batch(self, X, idx):
126
+ L = max(len(X[i]) for i in idx)
127
+ H = np.zeros((len(idx), L, X[idx[0]].shape[1]), dtype=np.float32)
128
+ mask = np.zeros((len(idx), L), dtype=bool)
129
+ for j, i in enumerate(idx):
130
+ H[j, :len(X[i])] = X[i]
131
+ mask[j, :len(X[i])] = True
132
+ return torch.from_numpy(H).to(self.device), torch.from_numpy(mask).to(self.device)
133
+
134
+ def _batches(self, n, shuffle):
135
+ order = np.random.permutation(n) if shuffle else np.arange(n)
136
+ for s in range(0, n, self.batch_size):
137
+ yield order[s:s + self.batch_size]
138
+
139
+ # ------------------------------------------------------------------ training
140
+ def fit(self, X, y=None):
141
+ self._set_seed()
142
+ X = self._check(X)
143
+ dim = X[0].shape[1]
144
+ self.net_ = _CVDDNet(dim, self.attention_size, self.n_heads).to(self.device)
145
+
146
+ # context vector initialisation (official initialize_context_vectors)
147
+ means = np.stack([x.mean(0) for x in X])
148
+ means = means / np.clip(np.linalg.norm(means, axis=1, keepdims=True), 1e-8, None)
149
+ km = KMeans(n_clusters=self.n_heads, n_init=10, random_state=self.random_state).fit(means)
150
+ centers = km.cluster_centers_ / np.linalg.norm(km.cluster_centers_, axis=1, keepdims=True)
151
+ self.net_.c.data = torch.from_numpy(centers.astype(np.float32)).to(self.device)
152
+
153
+ opt = torch.optim.Adam(self.net_.parameters(), lr=self.lr, weight_decay=self.weight_decay)
154
+ sched = torch.optim.lr_scheduler.MultiStepLR(opt, milestones=list(self.lr_milestones), gamma=0.1)
155
+ milestones, alphas = _alpha_schedule(self.alpha_scheduler, self.n_epochs)
156
+ alpha_i = 0
157
+ I = torch.eye(self.n_heads, device=self.device)
158
+
159
+ self.net_.alpha = 0.0
160
+ self.history_ = []
161
+ for epoch in range(self.n_epochs):
162
+ with warnings.catch_warnings(): # official order: step at the START of each epoch
163
+ warnings.simplefilter("ignore", UserWarning)
164
+ sched.step()
165
+ if epoch in milestones and alpha_i < len(alphas): # official: one step per epoch
166
+ self.net_.alpha = float(alphas[alpha_i])
167
+ alpha_i += 1
168
+ self.net_.train()
169
+ tot, nb = 0.0, 0
170
+ for idx in self._batches(len(X), shuffle=True):
171
+ H, mask = self._batch(X, idx)
172
+ d, w, _ = self.net_(H, mask)
173
+ P = torch.mean((self.net_.c @ self.net_.c.t() - I) ** 2)
174
+ loss = torch.mean(torch.sum(w * d, dim=1)) + self.lambda_p * P
175
+ opt.zero_grad()
176
+ loss.backward()
177
+ torch.nn.utils.clip_grad_norm_(self.net_.parameters(), 0.5) # as in the official trainer
178
+ opt.step()
179
+ tot += loss.item()
180
+ nb += 1
181
+ self.history_.append(tot / nb)
182
+ self._log(f"epoch {epoch + 1}/{self.n_epochs} loss {tot / nb:.6f} alpha {self.net_.alpha:g}")
183
+
184
+ return self._process_decision_scores(self.decision_function(X))
185
+
186
+ # ------------------------------------------------------------------ inference
187
+ @torch.no_grad()
188
+ def _forward_all(self, X):
189
+ X = self._check(X)
190
+ self.net_.eval()
191
+ D, A_all = [], []
192
+ for idx in self._batches(len(X), shuffle=False):
193
+ H, mask = self._batch(X, idx)
194
+ d, _, A = self.net_(H, mask)
195
+ D.append(d.cpu().numpy())
196
+ A = A.cpu().numpy()
197
+ A_all += [A[j, :, :len(X[i])] for j, i in enumerate(idx)]
198
+ return np.concatenate(D), A_all
199
+
200
+ def head_scores(self, X):
201
+ """Per-head cosine distances, shape [n_docs, n_heads]."""
202
+ return self._forward_all(X)[0]
203
+
204
+ def decision_function(self, X):
205
+ """Official 'context_dist_mean' score: mean cosine distance over heads."""
206
+ return self.head_scores(X).mean(1)
207
+
208
+ def attention(self, X):
209
+ """Attention weights per document, each [n_heads, n_tokens] (for inspection)."""
210
+ return self._forward_all(X)[1]
211
+
212
+ @torch.no_grad()
213
+ def token_scores(self, X):
214
+ """EXTENSION (not in the paper): per-token min_k 0.5 (1 - cos(h_t, c_k))."""
215
+ X = self._check(X)
216
+ C = F.normalize(self.net_.c.detach(), dim=1).cpu().numpy()
217
+ out = []
218
+ for x in X:
219
+ xn = x / np.clip(np.linalg.norm(x, axis=1, keepdims=True), 1e-8, None)
220
+ out.append((0.5 * (1 - xn @ C.T)).min(1))
221
+ return out