Skip to content
Draft
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
11 changes: 10 additions & 1 deletion chebai_graph/preprocessing/datasets/chebi.py
Original file line number Diff line number Diff line change
Expand Up @@ -168,7 +168,12 @@ def enc_if_not_none(encode, value):
assert len(encoded_values) == len(idents) == len(features)
torch.save(
[
{property.name: torch.cat(feat), "ident": id}
{
property.name: property.encoder.compress(
torch.cat(feat)
),
"ident": id,
}
Comment on lines 168 to +176
for feat, id in zip(encoded_values, idents)
if feat is not None
],
Expand Down Expand Up @@ -384,6 +389,8 @@ def load_processed_data(
property_data = torch.load(
self.get_property_path(property), weights_only=False
)
for entry in property_data:
entry[property.name] = property.encoder.decompress(entry[property.name])
if len(property_data[0][property.name].shape) > 1:
property.encoder.set_encoding_length(
property_data[0][property.name].shape[1]
Expand Down Expand Up @@ -535,6 +542,8 @@ def load_processed_data(
property_data = torch.load(
self.get_property_path(property), weights_only=False
)
for entry in property_data:
entry[property.name] = property.encoder.decompress(entry[property.name])
if len(property_data[0][property.name].shape) > 1:
property.encoder.set_encoding_length(
property_data[0][property.name].shape[1]
Expand Down
99 changes: 99 additions & 0 deletions chebai_graph/preprocessing/property_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,39 @@ def encode(self, value) -> torch.Tensor:
"""
return value

def compress(self, tensor: torch.Tensor) -> torch.Tensor:
"""
Compress an encoded tensor into a more compact on-disk representation.

Called just before caching property values to disk. The default
implementation is a no-op; subclasses override it to reduce file size
(e.g. by downcasting the dtype or storing indices instead of one-hot
vectors). Must be losslessly invertible by :meth:`decompress` (except
for deliberate float precision reductions).

Args:
tensor: The encoded property tensor for a single molecule.

Returns:
A compact tensor to store on disk.
"""
return tensor

def decompress(self, tensor: torch.Tensor) -> torch.Tensor:
"""
Reconstruct the full encoded tensor from its compressed on-disk form.

Inverse of :meth:`compress`, called right after loading cached property
values. The default implementation is a no-op.

Args:
tensor: The compressed property tensor as loaded from disk.

Returns:
The reconstructed encoded property tensor.
"""
return tensor

def on_start(self, **kwargs) -> None:
"""Hook called at the start of encoding process."""
pass
Expand Down Expand Up @@ -238,6 +271,62 @@ def encode(self, token: str | None) -> torch.Tensor:
self.tokens_dict[token], num_classes=self.get_encoding_length()
)

def compress(self, tensor: torch.Tensor) -> torch.Tensor:
"""
Store one index per node instead of the full one-hot matrix.

A dense ``(N, n_classes)`` int64 one-hot matrix is reduced to an
``(N,)`` vector of class indices. Index ``0`` is reserved for all-zero
rows (produced by :meth:`encode` for unknown tokens); real classes are
stored as ``argmax + 1``. The result uses ``uint8`` when it fits, else
``int16``.

Args:
tensor: One-hot tensor of shape ``(N, n_classes)``.

Returns:
Index tensor of shape ``(N,)``.
"""
if tensor.dim() != 2:
# already compressed / unexpected shape - leave untouched
return tensor
has_class = tensor.any(dim=1)
indices = torch.where(
has_class,
tensor.argmax(dim=1) + 1,
torch.zeros_like(has_class, dtype=torch.long),
)
dtype = torch.uint8 if tensor.shape[1] + 1 < 256 else torch.int16
return indices.to(dtype)
Comment on lines +299 to +300

def decompress(self, tensor: torch.Tensor) -> torch.Tensor:
"""
Reconstruct the dense one-hot matrix from stored class indices.

Inverse of :meth:`compress`. Index ``0`` maps back to an all-zero row;
index ``i > 0`` maps to a one-hot with class ``i - 1`` set.

Args:
tensor: Index tensor of shape ``(N,)`` as produced by
:meth:`compress`.

Returns:
One-hot tensor of shape ``(N, n_classes)`` keeping the compact
stored dtype (e.g. ``uint8``); it is promoted to float when merged
into the node/edge feature matrix at load time.
"""
if tensor.dim() != 1:
# already expanded / unexpected shape - leave untouched
return tensor
n_classes = self.get_encoding_length()
out = torch.zeros((tensor.shape[0], n_classes), dtype=tensor.dtype)
non_zero = tensor > 0
if non_zero.any():
out[non_zero] = torch.nn.functional.one_hot(
tensor[non_zero].to(torch.int64) - 1, num_classes=n_classes
).to(out.dtype)
return out


class AsIsEncoder(PropertyEncoder):
"""
Expand Down Expand Up @@ -271,6 +360,12 @@ def encode(self, token: float | int | None) -> torch.Tensor:
# ----- fix: for above warning
return torch.tensor(token).unsqueeze(0) # shape: (1, len(token))

def compress(self, tensor: torch.Tensor) -> torch.Tensor:
"""Downcast float values to float32 to halve the on-disk size."""
if tensor.is_floating_point():
return tensor.to(torch.float32)
return tensor


class BoolEncoder(PropertyEncoder):
"""
Expand All @@ -293,3 +388,7 @@ def encode(self, token: bool) -> torch.Tensor:
Tensor with 1 if True else 0.
"""
return torch.tensor([1 if token else 0])

def compress(self, tensor: torch.Tensor) -> torch.Tensor:
"""Store the 0/1 values as ``uint8`` instead of ``int64`` (8x smaller)."""
return tensor.to(torch.uint8)
Loading