kernel-elastic-autoencoder 1.0.0__tar.gz → 2.0.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.
Files changed (13) hide show
  1. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/PKG-INFO +1 -1
  2. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/pyproject.toml +4 -2
  3. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/losses.py +1 -1
  4. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/model.py +0 -44
  5. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/pipeline.py +5 -5
  6. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/sample.py +17 -5
  7. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/README.md +0 -0
  8. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/__init__.py +14 -14
  9. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/collate.py +0 -0
  10. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/config.py +0 -0
  11. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/layers.py +0 -0
  12. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/tokenizer.py +0 -0
  13. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/training.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: kernel_elastic_autoencoder
3
- Version: 1.0.0
3
+ Version: 2.0.0
4
4
  Summary: Implementation of Kernel-Elastic Autoencoder for Molecular Design (https://doi.org/10.1093/pnasnexus/pgae168)
5
5
  License: MIT
6
6
  Author: Felix Rotter-McCartney
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "kernel_elastic_autoencoder"
3
- version = "1.0.0"
3
+ version = "2.0.0"
4
4
  description = "Implementation of Kernel-Elastic Autoencoder for Molecular Design (https://doi.org/10.1093/pnasnexus/pgae168)"
5
5
  authors = [
6
6
  { name = "Felix Rotter-McCartney", email = "felix.rotter@mail.utoronto.ca" }
@@ -18,6 +18,9 @@ dependencies = [
18
18
  [tool.poetry]
19
19
  packages = [{ include = "kernel_elastic_autoencoder", from = "src" }]
20
20
 
21
+ [tool.ruff.lint]
22
+ ignore = ["B008"]
23
+
21
24
  [tool.semantic_release]
22
25
  version_toml = ["pyproject.toml:project.version"]
23
26
  commit_parser = "conventional"
@@ -40,7 +43,6 @@ dev = [
40
43
  "ruff (>=0.15.20,<0.17.0)",
41
44
  "pytest-cov (>=7.1.0,<8.0.0)",
42
45
  "pytest-xdist (>=3.8.0,<4.0.0)",
43
- "ty (>=0.0.62,<0.0.64)",
44
46
  "python-semantic-release (>=10.6.1,<11.0.0)",
45
47
  "pdoc (>=16.0.0,<17.0.0)",
46
48
  ]
@@ -121,7 +121,7 @@ class Loss(nn.Module):
121
121
  / (2 * (self.hp_sigma**2))
122
122
  ).sum()
123
123
  loss = self.hp_lambda * (
124
- torch.tensor([1])
124
+ torch.tensor(1, dtype=latents.dtype, device=latents.device)
125
125
  - ((1 / (latents.size(0) * self.kernel_dist_size)) * kernel_pairwise_sum)
126
126
  )
127
127
  return loss
@@ -1,9 +1,6 @@
1
- from pathlib import Path
2
- from typing import Any
3
1
 
4
2
  import torch
5
3
  from huggingface_hub import PyTorchModelHubMixin
6
- from huggingface_hub.hub_mixin import T, DataclassInstance
7
4
  from torch import nn
8
5
 
9
6
  from kernel_elastic_autoencoder.config import ModelConfig
@@ -78,47 +75,6 @@ class Model(
78
75
  )
79
76
  """Model decoder module."""
80
77
 
81
- @classmethod
82
- def from_pretrained(
83
- cls: type[T],
84
- pretrained_model_name_or_path: str | Path,
85
- *,
86
- force_download: bool = False,
87
- token: str | bool | None = None,
88
- cache_dir: str | Path | None = None,
89
- local_files_only: bool = False,
90
- revision: str | None = None,
91
- **model_kwargs,
92
- ) -> T:
93
- return cls.from_pretrained(
94
- pretrained_model_name_or_path,
95
- force_download=force_download,
96
- token=token,
97
- cache_dir=cache_dir,
98
- local_files_only=local_files_only,
99
- revision=revision,
100
- **model_kwargs,
101
- )
102
-
103
- def save_pretrained(
104
- self,
105
- save_directory: str | Path,
106
- *,
107
- config: dict | DataclassInstance | None = None,
108
- repo_id: str | None = None,
109
- push_to_hub: bool = False,
110
- model_card_kwargs: dict[str, Any] | None = None,
111
- **push_to_hub_kwargs,
112
- ) -> str | None:
113
- return super().save_pretrained(
114
- save_directory,
115
- config=config,
116
- repo_id=repo_id,
117
- push_to_hub=push_to_hub,
118
- model_card_kwargs=model_card_kwargs,
119
- **push_to_hub_kwargs,
120
- )
121
-
122
78
  def forward(
123
79
  self,
124
80
  input_ids: torch.Tensor,
@@ -31,7 +31,7 @@ class Completion(BaseModel):
31
31
  """Tensor containing generated condition masks for reuse."""
32
32
 
33
33
 
34
- class Pipeline[_T]:
34
+ class Pipeline[T]:
35
35
  """User-facing pipeline for inference.
36
36
 
37
37
  Defines an easy-to-use API for decoder-only inference with a pretrained model.
@@ -51,7 +51,7 @@ class Pipeline[_T]:
51
51
  self,
52
52
  model: Model,
53
53
  tokenizer: Tokenizer,
54
- collator: Callable[[ModelConfig, Tokenizer], Collator[_T]] = DataframeCollator,
54
+ collator: Callable[[ModelConfig, Tokenizer], Collator[T]] = DataframeCollator,
55
55
  sampler: Callable[[Tokenizer], Sampler] = Top1Sampler,
56
56
  device: torch.device | None = None,
57
57
  ) -> None:
@@ -87,7 +87,7 @@ class Pipeline[_T]:
87
87
  def _completion_entry(
88
88
  self,
89
89
  latents: torch.Tensor,
90
- dataset: _T,
90
+ dataset: T,
91
91
  seq_feature: str,
92
92
  cond_features: list[str],
93
93
  device: torch.device | None = None,
@@ -148,7 +148,7 @@ class Pipeline[_T]:
148
148
  self, intermediate: CompletionIntermediate, **kwargs
149
149
  ) -> Completion:
150
150
  return Completion(
151
- outputs=self.tokenizer.decode(
151
+ outputs=self.sampler(
152
152
  intermediate.input_ids, skip_special_tokens=True
153
153
  ),
154
154
  condition_embeddings=intermediate.condition_embeddings,
@@ -158,7 +158,7 @@ class Pipeline[_T]:
158
158
  def completion(
159
159
  self,
160
160
  latents: torch.Tensor,
161
- dataset: _T,
161
+ dataset: T,
162
162
  seq_feature: str,
163
163
  cond_features: list[str],
164
164
  device: torch.device | None = None,
@@ -12,25 +12,34 @@ class Sampler(ABC):
12
12
  Provides a specification for sampling text sequences from logits.
13
13
  """
14
14
 
15
- def __init__(self, tokenizer: Tokenizer):
15
+ def __init__(self, tokenizer: Tokenizer, **kwargs):
16
16
  """Instantiates a Sampler object.
17
17
 
18
18
  Args:
19
19
  tokenizer: Object implementing the Tokenizer protocol.
20
+ **kwargs: Sampler-specific keyword args.
20
21
  """
21
22
  self.tokenizer = tokenizer
22
23
  """Object implementing the Tokenizer protocol."""
23
24
 
25
+ @abstractmethod
24
26
  def __call__(
25
27
  self, logits: torch.Tensor, skip_special_tokens: bool, **kwargs
26
28
  ) -> Iterable[str]:
27
- ids = self.sample_ids(logits, **kwargs)
28
- sampled = self.tokenizer.decode(ids, skip_special_tokens=skip_special_tokens)
29
- return sampled
29
+ """Interface method for implementing the last sampling step where outputs are decoded through the tokenizer.
30
+
31
+ Args:
32
+ logits: Tensor of dimension (B, S) containing IDs at the last step of inference.
33
+ **kwargs: Keyword arguments.
34
+
35
+ Returns:
36
+ Iterable[str]: Iterable of strings decoded by the tokenizer.
37
+ """
38
+ ...
30
39
 
31
40
  @abstractmethod
32
41
  def sample_ids(self, logits: torch.Tensor, **kwargs) -> torch.Tensor:
33
- """Interface method for implementing index sampling from logits.
42
+ """Interface method for implementing index sampling from logits in intermediate steps.
34
43
 
35
44
  Args:
36
45
  logits: Tensor of dimension (B, S, L) containing logits.
@@ -43,6 +52,9 @@ class Sampler(ABC):
43
52
 
44
53
 
45
54
  class Top1Sampler(Sampler):
55
+ def __call__(self, logits: torch.Tensor, skip_special_tokens: bool, **kwargs) -> Iterable[str]:
56
+ return self.tokenizer.decode(logits, skip_special_tokens=skip_special_tokens)
57
+
46
58
  def sample_ids(self, logits: torch.Tensor, **kwargs) -> torch.Tensor:
47
59
  """Implementation of index sampling from logits choosing the highest-probability token.
48
60
 
@@ -19,26 +19,26 @@ from kernel_elastic_autoencoder.tokenizer import Tokenizer
19
19
  from kernel_elastic_autoencoder.training import Trainer, TrainerCallback
20
20
 
21
21
  __all__ = [
22
+ "Collated",
23
+ "Collator",
24
+ "Completion",
25
+ "DataframeCollator",
22
26
  "ExperimentConfig",
23
- "ModelConfig",
27
+ "Loss",
28
+ "Model",
24
29
  "ModelCommonConfig",
25
- "ModelInputConfig",
26
- "ModelEncoderConfig",
30
+ "ModelConfig",
27
31
  "ModelDecoderConfig",
28
- "TrainingConfig",
29
- "TrainingCommonConfig",
30
- "TrainingHyperparameterConfig",
31
- "TrainingOptimizerConfig",
32
+ "ModelEncoderConfig",
33
+ "ModelInputConfig",
32
34
  "Pipeline",
33
- "Completion",
34
- "Model",
35
- "Loss",
36
- "Tokenizer",
37
- "Collator",
38
- "Collated",
39
- "DataframeCollator",
40
35
  "Sampler",
36
+ "Tokenizer",
41
37
  "Top1Sampler",
42
38
  "Trainer",
43
39
  "TrainerCallback",
40
+ "TrainingCommonConfig",
41
+ "TrainingConfig",
42
+ "TrainingHyperparameterConfig",
43
+ "TrainingOptimizerConfig",
44
44
  ]