Skip to content

Commit 03bcf02

Browse files
Merge pull request #131 from x-tabdeveloping/latent_terms
[BETA] Added a Latent terms implementation
2 parents 149d878 + 659f6bf commit 03bcf02

8 files changed

Lines changed: 290 additions & 7 deletions

File tree

pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@ profile = "black"
99

1010
[project]
1111
name = "turftopic"
12-
version = "0.25.2"
12+
version = "0.25.3"
1313
description = "Topic modeling with contextual representations from sentence transformers."
1414
authors = [
1515
{ name = "Márton Kardos <power.up1163@gmail.com>", email = "martonkardos@cas.au.dk" }

turftopic/late.py

Lines changed: 15 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,10 @@
11
import itertools
22
import warnings
3+
from functools import partial
34
from typing import Callable, Iterable, Optional, Union
45

56
import numpy as np
7+
import scipy.sparse as spr
68
import torch
79
from sentence_transformers import SentenceTransformer
810
from sklearn.base import TransformerMixin
@@ -208,29 +210,37 @@ def unflatten_repr(
208210
return repr
209211

210212

211-
def pool_flat(flat_repr: np.ndarray, lengths: Lengths, agg=np.nanmean):
213+
def pool_flat(
214+
flat_repr: np.ndarray | spr.sparray, lengths: Lengths, agg=np.nanmean
215+
):
212216
"""Pools vectors within documents using the agg function.
213217
214218
Parameters
215219
----------
216-
flat_repr: ndarray of shape (n_total_tokens, n_dims)
220+
flat_repr: ndarray or sparse array of shape (n_total_tokens, n_dims)
217221
Flattened document representations.
218222
lengths: Lengths
219223
Number of tokens in each document.
220224
221225
Returns
222226
-------
223-
ndarray of shape (n_documents, n_dims)
227+
ndarray or sparse array of shape (n_documents, n_dims)
224228
Pooled representation for each document.
225229
"""
230+
if spr.issparse(flat_repr):
231+
stack = partial(spr.vstack, format="csr")
232+
array = spr.csr_matrix
233+
else:
234+
stack = np.stack
235+
array = np.asarray
226236
pooled = []
227237
start_index = 0
228238
for length in lengths:
229239
pooled.append(
230-
agg(flat_repr[start_index : start_index + length], axis=0)
240+
array(agg(flat_repr[start_index : start_index + length], axis=0))
231241
)
232242
start_index += length
233-
return np.stack(pooled)
243+
return stack(pooled)
234244

235245

236246
def get_document_chunks(

turftopic/retrieval/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
from .bm25 import BM25Transformer

turftopic/retrieval/bm25.py

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,32 @@
1+
import numpy as np
2+
import scipy.sparse as spr
3+
from sklearn.base import BaseEstimator, TransformerMixin
4+
5+
6+
class BM25Transformer(BaseEstimator, TransformerMixin):
7+
def __init__(self, b: float = 0.7, k1: float = 8):
8+
self.b = b
9+
self.k1 = k1
10+
11+
def fit(self, X, y=None):
12+
self.N_ = X.shape[0]
13+
self.avgdl_ = X.sum(axis=1).mean()
14+
self.term_freq_ = np.ravel(np.asarray((X > 0).sum(axis=0)))
15+
self.idf_ = np.log(
16+
(self.N_ - self.term_freq_ + 0.5) / (self.term_freq_ + 0.5)
17+
)
18+
return self
19+
20+
def transform(self, X):
21+
if spr.issparse(X):
22+
X = spr.csr_array(X)
23+
d_len = np.ravel(np.asarray(X.sum(axis=1)))
24+
K_D = 1 - self.b + self.b * d_len / self.avgdl_
25+
return (
26+
self.idf_[None, :]
27+
* (X * (self.k1 + 1))
28+
/ (X + self.k1 * K_D[:, None])
29+
)
30+
31+
def fit_transform(self, X, y=None):
32+
return self.fit(X, y).transform(X)

turftopic/serialization.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -77,7 +77,11 @@ def validate_package_versions(remote_versions: dict[str, str]):
7777

7878
def create_readme(model, model_path: str) -> str:
7979
model_structure = str(model)
80-
topics_table = model.export_topics(format="markdown", top_k=10)
80+
try:
81+
topics_table = model.export_topics(format="markdown", top_k=10)
82+
except Exception:
83+
print("Couldn't produce topic table for readme, moving on...")
84+
topics_table = None
8185
local_versions = get_package_versions()
8286
lines = ["| Package | Version |", "| - | - |"]
8387
for package in IMPORTANT_PACKAGES:

turftopic/vectorizers/__init__.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
from turftopic.vectorizers.latent_terms.latent_terms import (
2+
LatentTermsVectorizer,
3+
)
Lines changed: 112 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,112 @@
1+
import json
2+
import tempfile
3+
from pathlib import Path
4+
from typing import Union
5+
6+
import joblib
7+
import numpy as np
8+
from huggingface_hub import HfApi
9+
from sklearn.base import BaseEstimator, TransformerMixin
10+
11+
from turftopic.late import (
12+
LateSentenceTransformer,
13+
flatten_repr,
14+
pool_flat,
15+
)
16+
from turftopic.serialization import create_readme, get_package_versions
17+
from turftopic.vectorizers.latent_terms.top_k_autoencoder import (
18+
TopKAutoEncoder,
19+
)
20+
21+
22+
class LatentTermsVectorizer(BaseEstimator, TransformerMixin):
23+
def __init__(
24+
self,
25+
encoder: str | LateSentenceTransformer,
26+
autoencoder: TopKAutoEncoder,
27+
concept_labels: np.ndarray,
28+
show_progress_bar: bool = True,
29+
):
30+
self.encoder = encoder
31+
if isinstance(self.encoder, str):
32+
self._encoder = LateSentenceTransformer(self.encoder)
33+
else:
34+
self._encoder = self.encoder
35+
self.concept_labels = np.array(concept_labels)
36+
self.autoencoder = autoencoder
37+
self.show_progress_bar = show_progress_bar
38+
self.autoencoder.show_progress_bar = show_progress_bar
39+
40+
def fit(self, raw_documents, y=None):
41+
# Does nothing, for compatibility
42+
return self
43+
44+
def transform(self, raw_documents):
45+
token_embeddings, offsets = self._encoder.encode_tokens(
46+
list(raw_documents), show_progress_bar=self.show_progress_bar
47+
)
48+
flat_token_embeddings, lengths = flatten_repr(token_embeddings)
49+
flat_z = self.autoencoder.transform(flat_token_embeddings)
50+
# Pooling procedure from section 3.2
51+
pooled_z = pool_flat(flat_z, lengths=lengths, agg=np.sum)
52+
return np.sqrt(pooled_z)
53+
54+
def fit_transform(self, raw_documents, y=None):
55+
return self.fit(raw_documents, y).transform(raw_documents)
56+
57+
def get_feature_names_out(self):
58+
return self.concept_labels
59+
60+
@classmethod
61+
def from_dict(cls, data):
62+
autoencoder = TopKAutoEncoder.from_dict(data["autoencoder"])
63+
return cls(
64+
encoder=data["encoder"],
65+
autoencoder=autoencoder,
66+
show_progress_bar=data["show_progress_bar"],
67+
concept_labels=data["concept_labels"],
68+
)
69+
70+
def to_dict(self):
71+
return dict(
72+
encoder=self.encoder,
73+
autoencoder=self.autoencoder.to_dict(),
74+
show_progress_bar=self.show_progress_bar,
75+
concept_labels=self.concept_labels,
76+
)
77+
78+
def to_disk(self, out_dir: Union[Path, str]):
79+
"""Persists model to directory on your machine.
80+
81+
Parameters
82+
----------
83+
out_dir: Path | str
84+
Directory to save the model to.
85+
"""
86+
out_dir = Path(out_dir)
87+
out_dir.mkdir(exist_ok=True)
88+
package_versions = get_package_versions()
89+
with out_dir.joinpath("package_versions.json").open("w") as ver_file:
90+
ver_file.write(json.dumps(package_versions))
91+
joblib.dump(self, out_dir.joinpath("model.joblib"))
92+
93+
def push_to_hub(self, repo_id: str):
94+
"""Uploads model to HuggingFace Hub
95+
96+
Parameters
97+
----------
98+
repo_id: str
99+
Repository to upload the model to.
100+
"""
101+
api = HfApi()
102+
api.create_repo(repo_id, exist_ok=True)
103+
with tempfile.TemporaryDirectory() as tmp_dir:
104+
readme_path = Path(tmp_dir).joinpath("README.md")
105+
with readme_path.open("w") as readme_file:
106+
readme_file.write(create_readme(self, repo_id))
107+
self.to_disk(tmp_dir)
108+
api.upload_folder(
109+
folder_path=tmp_dir,
110+
repo_id=repo_id,
111+
repo_type="model",
112+
)
Lines changed: 121 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,121 @@
1+
"""This is an encode-only implementation of the TopK autoencoder.
2+
The training code lives in the x-tabdeveloping/latent_terms GitHub repo"""
3+
4+
import warnings
5+
from functools import partial
6+
from typing import Optional
7+
8+
import numpy as np
9+
import scipy.sparse as spr
10+
from sklearn.base import BaseEstimator, TransformerMixin
11+
from tqdm import trange
12+
13+
try:
14+
import jax.numpy as jnp
15+
from jax import jit
16+
from jax.lax import top_k
17+
except ModuleNotFoundError:
18+
warnings.warn("JAX not found, continuing with NumPy implementation.")
19+
jnp = np
20+
21+
# Dummy JIT as the identity function
22+
def jit(f):
23+
return f
24+
25+
# NumPy implementation of the TopK activation function.
26+
def top_k(a, k, *, axis=-1):
27+
if axis is None:
28+
axis_size = a.size
29+
else:
30+
axis_size = a.shape[axis]
31+
index_array = np.argpartition(a, axis_size - k, axis=axis)
32+
topk_indices = np.take(index_array, -np.arange(k) - 1, axis=axis)
33+
topk_values = np.take_along_axis(a, topk_indices, axis=axis)
34+
return topk_values, topk_indices
35+
36+
37+
def top_k_activation(z, k: int):
38+
values, indices = top_k(z, k=k, axis=-1)
39+
threshold = jnp.min(values, axis=-1)
40+
condition = threshold[:, None] <= z
41+
return jnp.where(condition, z, 0)
42+
43+
44+
def encode(params, x, k: int):
45+
z = x @ params["W_e"] + params["b_e"]
46+
return top_k_activation(z, k)
47+
48+
49+
class TopKAutoEncoder(BaseEstimator, TransformerMixin):
50+
def __init__(
51+
self,
52+
n_latent: int = 32768,
53+
top_k: int = 16,
54+
lr: float = 1e-3,
55+
batch_size: int = 4096,
56+
n_epochs: int = 10,
57+
alpha: float = 0.03,
58+
show_progress_bar: bool = True,
59+
random_state: Optional[int] = None,
60+
):
61+
self.random_state = random_state
62+
self.n_latent = n_latent
63+
self.lr = lr
64+
self.alpha = alpha
65+
self.top_k = top_k
66+
self.batch_size = batch_size
67+
self.n_epochs = n_epochs
68+
self.show_progress_bar = show_progress_bar
69+
70+
def fit(self, X, y=None):
71+
# Training is implemented here: https://github.com/x-tabdeveloping/latent_terms
72+
return self
73+
74+
def to_dict(self) -> dict:
75+
return dict(
76+
attr=self.get_params(),
77+
params=self._params,
78+
loss_curve=self.loss_curve_,
79+
)
80+
81+
@classmethod
82+
def from_dict(cls, data):
83+
obj = cls(**data["attr"])
84+
params = data["params"]
85+
obj.coef_ = np.array(params["W_e"])
86+
obj.coef_d_ = np.array(params["W_d"])
87+
obj.intercept_ = np.array(params["b_e"])
88+
obj.intercept_d_ = np.array(params["b_d"])
89+
obj.loss_curve_ = data["loss_curve"]
90+
return obj
91+
92+
@property
93+
def _params(self):
94+
return {
95+
"W_e": self.coef_,
96+
"b_e": self.intercept_,
97+
"W_d": self.coef_d_,
98+
"b_d": self.intercept_d_,
99+
}
100+
101+
def transform(self, X):
102+
if spr.issparse(X):
103+
X = X.todense()
104+
Z = []
105+
_encode = jit(partial(encode, params=self._params, k=self.top_k))
106+
for batch_start in trange(
107+
0,
108+
X.shape[0],
109+
self.batch_size,
110+
leave=False,
111+
desc="Going through all batches",
112+
disable=not self.show_progress_bar,
113+
):
114+
batch_end = batch_start + self.batch_size
115+
batch_x = X[batch_start:batch_end]
116+
batch_z = _encode(x=batch_x)
117+
Z.append(spr.csr_array(batch_z))
118+
return spr.vstack(Z, format="csr")
119+
120+
def fit_transform(self, X, y=None):
121+
return self.fit(X, y).transform(X)

0 commit comments

Comments
 (0)