metalsurfer 0.3.0__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.
Files changed (41) hide show
  1. metalsurfer/__init__.py +147 -0
  2. metalsurfer/_logging.py +325 -0
  3. metalsurfer/_utils.py +13 -0
  4. metalsurfer/campaigns.py +583 -0
  5. metalsurfer/config.py +695 -0
  6. metalsurfer/conformers.py +249 -0
  7. metalsurfer/exceptions.py +29 -0
  8. metalsurfer/filters.py +604 -0
  9. metalsurfer/io_results.py +751 -0
  10. metalsurfer/ml/__init__.py +11 -0
  11. metalsurfer/ml/bayesian.py +927 -0
  12. metalsurfer/ml/dataset.py +217 -0
  13. metalsurfer/ml/features.py +147 -0
  14. metalsurfer/ml/predict.py +123 -0
  15. metalsurfer/ml/regression.py +307 -0
  16. metalsurfer/ml/reproduce.py +158 -0
  17. metalsurfer/ml/schema.py +579 -0
  18. metalsurfer/models.py +773 -0
  19. metalsurfer/optimization.py +1300 -0
  20. metalsurfer/placement/__init__.py +64 -0
  21. metalsurfer/placement/_constants.py +211 -0
  22. metalsurfer/placement/_material.py +59 -0
  23. metalsurfer/placement/generators.py +1435 -0
  24. metalsurfer/placement/geometry.py +856 -0
  25. metalsurfer/placement/policy.py +206 -0
  26. metalsurfer/placement/sites.py +1815 -0
  27. metalsurfer/py.typed +0 -0
  28. metalsurfer/surface_prep/__init__.py +69 -0
  29. metalsurfer/surface_prep/prep.py +372 -0
  30. metalsurfer/surfaces.py +1004 -0
  31. metalsurfer/symmetry.py +460 -0
  32. metalsurfer/workflow/__init__.py +15 -0
  33. metalsurfer/workflow/bayesian.py +601 -0
  34. metalsurfer/workflow/core.py +498 -0
  35. metalsurfer/workflow/reference.py +92 -0
  36. metalsurfer/workflow/saturation.py +913 -0
  37. metalsurfer/workflow/shared.py +1070 -0
  38. metalsurfer-0.3.0.dist-info/METADATA +41 -0
  39. metalsurfer-0.3.0.dist-info/RECORD +41 -0
  40. metalsurfer-0.3.0.dist-info/WHEEL +5 -0
  41. metalsurfer-0.3.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,147 @@
1
+ """Adsorption on arbitrary materials."""
2
+
3
+ __version__ = "0.3.0"
4
+
5
+ import importlib
6
+
7
+ from ._logging import ensure_log_record_defaults
8
+ from .config import AdsorptionConfig
9
+ from .exceptions import (
10
+ DependencyMissingError,
11
+ GeometryValidationError,
12
+ OptimizationError,
13
+ )
14
+ from .models import (
15
+ BindingCampaignResult,
16
+ MoleculeCampaignSummary,
17
+ MoleculeSummary,
18
+ MultiMolSaturationRunResult,
19
+ MultiMolSaturationStepResult,
20
+ ReferenceEnergies,
21
+ SaturationCampaignResult,
22
+ SaturationRunResult,
23
+ SaturationStepResult,
24
+ ScreeningResult,
25
+ ScreeningRunResult,
26
+ TimingInfo,
27
+ )
28
+
29
+ __all__ = [
30
+ "__version__",
31
+ "AdsorptionConfig",
32
+ "run_adsorption",
33
+ "run_adsorption_bo",
34
+ "run_saturation",
35
+ "run_saturation_bo",
36
+ "BindingCampaignResult",
37
+ "MoleculeCampaignSummary",
38
+ "MoleculeSummary",
39
+ "ReferenceEnergies",
40
+ "SaturationCampaignResult",
41
+ "SaturationRunResult",
42
+ "SaturationStepResult",
43
+ "MultiMolSaturationRunResult",
44
+ "MultiMolSaturationStepResult",
45
+ "ScreeningResult",
46
+ "ScreeningRunResult",
47
+ "TimingInfo",
48
+ "DependencyMissingError",
49
+ "GeometryValidationError",
50
+ "OptimizationError",
51
+ "configure_logging",
52
+ ]
53
+
54
+ _LAZY_MODULES = {
55
+ "_logging": {"configure_logging"},
56
+ "surface_prep": {
57
+ "prepare_substrate",
58
+ "finalize_substrate",
59
+ "relax_substrate",
60
+ "resize_substrate_for_molecule",
61
+ "apply_material_pbc",
62
+ "SlabContainer",
63
+ "create_slab_from_bulk",
64
+ "create_slab_from_atoms",
65
+ "substitute_alloy",
66
+ "deposit_adatoms",
67
+ "auto_resize_substrate_for_molecule",
68
+ "compute_minimum_supercell",
69
+ "ensure_slab_z_alignment",
70
+ "apply_surface_constraints",
71
+ "validate_substrate",
72
+ "accept_substrate_for_api",
73
+ "coerce_slab_container",
74
+ },
75
+ "conformers": {"create_conformers_from_smiles", "select_conformer_boltzmann"},
76
+ "placement": {
77
+ "generate_placement_from_spec",
78
+ "generate_placement_from_descriptor",
79
+ "enumerate_placement_specs",
80
+ "calculate_min_distance",
81
+ "get_symmetry_aware_sites",
82
+ "get_symmetry_info",
83
+ },
84
+ "optimization": {
85
+ "setup_calculator",
86
+ "setup_torchsim_model",
87
+ "setup_single_model",
88
+ "TorchSimCalculator",
89
+ "optimize_isolated_molecules_batched",
90
+ "optimize_adsorbate_slab_batched",
91
+ "batch_static",
92
+ "identify_top_layer_indices",
93
+ "identify_relaxable_surface_indices",
94
+ "compute_frozen_indices",
95
+ "frozen_indices_from_constraints",
96
+ },
97
+ "filters": {"filter_results", "check_decomposition", "check_desorption"},
98
+ "workflow": {
99
+ "process_molecule",
100
+ "process_molecule_bayesian",
101
+ "run_saturation_screening",
102
+ "calculate_reference_energies",
103
+ "load_molecules",
104
+ },
105
+ "campaigns": {
106
+ "run_adsorption",
107
+ "run_adsorption_bo",
108
+ "run_saturation",
109
+ "run_saturation_bo",
110
+ },
111
+ "io_results": {
112
+ "setup_directories",
113
+ "save_molecule_results",
114
+ "save_single_molecule_results",
115
+ "screening_run_result",
116
+ "save_summary_results",
117
+ "save_saturation_results",
118
+ "save_multi_mol_saturation_results",
119
+ "write_run_metadata",
120
+ "write_run_settings",
121
+ "results_dir",
122
+ "results_dir_for",
123
+ },
124
+ "ml": {
125
+ "BindingEnergyPredictor",
126
+ "ComputationContext",
127
+ "DatasetLogger",
128
+ "PlacementRecord",
129
+ "evaluate_model",
130
+ "extract_features",
131
+ "extract_features_from_dataset",
132
+ "grouped_cross_validate",
133
+ "load_dataset",
134
+ "train_model",
135
+ },
136
+ "symmetry": {"SymmetryAnalyzer", "SymmetryAnalysisError"},
137
+ }
138
+
139
+ ensure_log_record_defaults()
140
+
141
+
142
+ def __getattr__(name: str):
143
+ # Intentional lazy import: heavy/optional submodules loaded on first access.
144
+ for mod, names in _LAZY_MODULES.items():
145
+ if name in names:
146
+ return getattr(importlib.import_module(f".{mod}", __name__), name)
147
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
@@ -0,0 +1,325 @@
1
+ """Context vars and formatters for structured logging (e.g. ``ctx_prefix``)."""
2
+
3
+ import contextvars
4
+ import io
5
+ import logging
6
+ import os
7
+ import sys
8
+ import time
9
+ from contextlib import contextmanager
10
+ from typing import Any
11
+
12
+ _LOG_CTX: contextvars.ContextVar[dict[str, Any] | None] = contextvars.ContextVar(
13
+ "adsorption_log_ctx", default=None
14
+ )
15
+ _FACTORY_INSTALLED = False
16
+ CTX_KEY_ORDER = ("molecule", "surface_type", "placement_id", "seed")
17
+
18
+
19
+ @contextmanager
20
+ def log_context(**kwargs: Any):
21
+ """Push key-value pairs into logging context for this scope."""
22
+ prev = _LOG_CTX.get() or {}
23
+ merged = {**prev, **kwargs}
24
+ token = _LOG_CTX.set(merged)
25
+ try:
26
+ yield merged
27
+ finally:
28
+ _LOG_CTX.reset(token)
29
+
30
+
31
+ def get_log_context() -> dict[str, Any]:
32
+ """Return current logging context (read-only)."""
33
+ return dict(_LOG_CTX.get() or {})
34
+
35
+
36
+ _warned_once: set[str] = set()
37
+
38
+
39
+ def warn_once(logger: logging.Logger, key: str, message: str) -> None:
40
+ """Emit a warning log message at most once per *key* across the process."""
41
+ if key not in _warned_once:
42
+ _warned_once.add(key)
43
+ logger.warning(message)
44
+
45
+
46
+ def _format_ctx_prefix(ctx: dict[str, Any]) -> str:
47
+ if not ctx:
48
+ return ""
49
+ parts = []
50
+ for k in CTX_KEY_ORDER:
51
+ if k in ctx:
52
+ parts.append(f"{k}={ctx[k]}")
53
+ for k, v in ctx.items():
54
+ if k not in CTX_KEY_ORDER:
55
+ parts.append(f"{k}={v}")
56
+ return "[" + " ".join(parts) + "] "
57
+
58
+
59
+ class ContextFilter(logging.Filter):
60
+ """Inject ctx_prefix into log records from current context."""
61
+
62
+ _KEY_ORDER = CTX_KEY_ORDER
63
+
64
+ def filter(self, record: logging.LogRecord) -> bool:
65
+ ctx = _LOG_CTX.get() or {}
66
+ record.ctx_prefix = _format_ctx_prefix(ctx)
67
+ return True
68
+
69
+
70
+ def _ctx_prefix_from_context() -> str:
71
+ return _format_ctx_prefix(_LOG_CTX.get() or {})
72
+
73
+
74
+ def _install_log_record_factory() -> None:
75
+ """Ensure every LogRecord always has ctx_prefix."""
76
+ global _FACTORY_INSTALLED
77
+ if _FACTORY_INSTALLED:
78
+ return
79
+
80
+ old_factory = logging.getLogRecordFactory()
81
+
82
+ def record_factory(*args: Any, **kwargs: Any) -> logging.LogRecord:
83
+ record = old_factory(*args, **kwargs)
84
+ if not hasattr(record, "ctx_prefix"):
85
+ record.ctx_prefix = _ctx_prefix_from_context()
86
+ return record
87
+
88
+ logging.setLogRecordFactory(record_factory)
89
+ _FACTORY_INSTALLED = True
90
+
91
+
92
+ def ensure_log_record_defaults() -> None:
93
+ """Ensure LogRecords include required contextual fields."""
94
+ _install_log_record_factory()
95
+
96
+
97
+ class _LogStreamToLogger(io.TextIOBase):
98
+ """File-like object that forwards writes to a logger.
99
+
100
+ Designed for capturing non-logging output (e.g. progress bars) and turning
101
+ it into coherent line-based log records.
102
+ """
103
+
104
+ def __init__(
105
+ self,
106
+ *,
107
+ logger: logging.Logger,
108
+ level: int,
109
+ carriage_return_rate_limit_s: float = 1.0,
110
+ ):
111
+ super().__init__()
112
+ self._logger = logger
113
+ self._level = level
114
+ self._rate_limit_s = float(carriage_return_rate_limit_s)
115
+
116
+ self._pending: str = ""
117
+ self._last_cr_text: str = ""
118
+ self._last_emit_t: float = time.monotonic()
119
+ self._last_emit_msg: str = ""
120
+
121
+ def writable(self) -> bool: # pragma: no cover
122
+ return True
123
+
124
+ def isatty(self) -> bool: # pragma: no cover
125
+ return False
126
+
127
+ def _maybe_emit_cr_snapshot(self, snapshot: str) -> None:
128
+ msg = snapshot.strip()
129
+ if not msg:
130
+ return
131
+
132
+ now = time.monotonic()
133
+ if now - self._last_emit_t >= self._rate_limit_s:
134
+ if msg != self._last_emit_msg:
135
+ self._logger.log(self._level, msg)
136
+ self._last_emit_t = now
137
+ self._last_emit_msg = msg
138
+ self._last_cr_text = ""
139
+
140
+ def _emit_final_line(self) -> None:
141
+ msg = self._pending.strip()
142
+ self._pending = ""
143
+ if msg:
144
+ self._logger.log(self._level, msg)
145
+ self._last_emit_t = time.monotonic()
146
+ self._last_emit_msg = msg
147
+ self._last_cr_text = ""
148
+
149
+ def flush(self) -> None:
150
+ # Emit any trailing content at context exit.
151
+ if self._pending.strip():
152
+ self._emit_final_line()
153
+ return
154
+ if self._last_cr_text.strip():
155
+ self._logger.log(self._level, self._last_cr_text.strip())
156
+ self._last_cr_text = ""
157
+
158
+ def write(self, s: str) -> int:
159
+ if not s:
160
+ return 0
161
+ if not isinstance(s, str):
162
+ s = str(s)
163
+
164
+ for ch in s:
165
+ if ch == "\n":
166
+ if self._pending:
167
+ self._emit_final_line()
168
+ else:
169
+ if self._last_cr_text.strip():
170
+ self._logger.log(self._level, self._last_cr_text.strip())
171
+ self._last_cr_text = ""
172
+ continue
173
+
174
+ if ch == "\r":
175
+ snapshot = self._pending
176
+ self._pending = ""
177
+ if snapshot.strip():
178
+ self._last_cr_text = snapshot.strip()
179
+ self._maybe_emit_cr_snapshot(snapshot)
180
+ else:
181
+ self._last_cr_text = ""
182
+ continue
183
+
184
+ self._pending += ch
185
+
186
+ return len(s)
187
+
188
+
189
+ def _parse_level(level_name: str, default: int) -> int:
190
+ level = getattr(logging, str(level_name).upper(), None)
191
+ return level if isinstance(level, int) else default
192
+
193
+
194
+ def _ensure_context_filter(handler: logging.Handler) -> None:
195
+ """Attach ContextFilter once so ctx_prefix is set at format time."""
196
+ if not any(isinstance(f, ContextFilter) for f in handler.filters):
197
+ handler.addFilter(ContextFilter())
198
+
199
+
200
+ def configure_logging(
201
+ *,
202
+ default_level: str = "INFO",
203
+ fmt: str = "%(asctime)s %(name)s %(levelname)s %(ctx_prefix)s%(message)s",
204
+ datefmt: str = "%H:%M:%S",
205
+ ) -> None:
206
+ """Configure project logging with sane HPC defaults.
207
+
208
+ Environment overrides:
209
+ - METALSURFER_LOG_LEVEL (default: INFO)
210
+ - TORCHSIM_LOG_LEVEL (default: WARNING)
211
+
212
+ Notes:
213
+ - Metalsurfer logs are routed to stdout (not stderr) so HPC job
214
+ launchers that split stdout/stderr capture INFO logs in `.out`.
215
+ - When running under pytest, we avoid reconfiguring stream handlers by
216
+ default to avoid interacting badly with pytest's logging capture.
217
+ """
218
+ _install_log_record_factory()
219
+
220
+ root = logging.getLogger()
221
+ level_name = os.getenv("METALSURFER_LOG_LEVEL", default_level)
222
+ formatter = logging.Formatter(fmt=fmt, datefmt=datefmt)
223
+ running_pytest = "PYTEST_CURRENT_TEST" in os.environ or any(
224
+ name.startswith("pytest") for name in sys.modules
225
+ )
226
+ force_stdout = os.getenv("METALSURFER_FORCE_STDOUT_LOGS", "") == "1"
227
+ target_stream = sys.stdout if (force_stdout or not running_pytest) else sys.stderr
228
+
229
+ root.setLevel(_parse_level(level_name, root.level))
230
+
231
+ # Avoid touching stream handlers under pytest unless explicitly forced.
232
+ if not (running_pytest and not force_stdout):
233
+ if not root.handlers:
234
+ logging.basicConfig(
235
+ level=_parse_level(level_name, logging.INFO),
236
+ format=fmt,
237
+ datefmt=datefmt,
238
+ stream=target_stream,
239
+ )
240
+ else:
241
+ root.setLevel(_parse_level(level_name, root.level))
242
+
243
+ stream_handler_found = False
244
+ for handler in root.handlers:
245
+ if isinstance(handler, logging.StreamHandler):
246
+ handler.setStream(target_stream)
247
+ handler.setFormatter(formatter)
248
+ _ensure_context_filter(handler)
249
+ stream_handler_found = True
250
+
251
+ if not stream_handler_found:
252
+ sh = logging.StreamHandler(target_stream)
253
+ sh.setFormatter(formatter)
254
+ _ensure_context_filter(sh)
255
+ root.addHandler(sh)
256
+
257
+ torchsim_level = _parse_level(
258
+ os.getenv("TORCHSIM_LOG_LEVEL", "WARNING"),
259
+ logging.WARNING,
260
+ )
261
+ for logger_name in ("torch_sim", "torchsim", "torch_sim.autobatching"):
262
+ torch_logger = logging.getLogger(logger_name)
263
+ torch_logger.setLevel(torchsim_level)
264
+ # Avoid changing propagation if TorchSim already installed its own handlers.
265
+ # When there are no handlers, propagation ensures logs reach our root setup.
266
+ if not torch_logger.handlers:
267
+ torch_logger.propagate = True
268
+ if not (running_pytest and not force_stdout):
269
+ for handler in torch_logger.handlers:
270
+ if isinstance(handler, logging.StreamHandler):
271
+ handler.setStream(target_stream)
272
+ handler.setFormatter(formatter)
273
+
274
+
275
+ @contextmanager
276
+ def torchsim_output_capture(
277
+ *,
278
+ logger_name: str = "metalsurfer.torchsim",
279
+ stdout_level: int = logging.INFO,
280
+ stderr_level: int = logging.WARNING,
281
+ carriage_return_rate_limit_s: float = 1.0,
282
+ ):
283
+ """Capture TorchSim's stdout/stderr and route through logging.
284
+
285
+ Useful because TorchSim (and its progress bars) may print directly to
286
+ stdout/stderr, bypassing Python's logging configuration.
287
+
288
+ Notes:
289
+ - stdout is mapped to INFO, stderr is mapped to WARNING.
290
+ - stdout updates using carriage return (``\\r``) are rate-limited so we
291
+ don't emit thousands of near-identical log lines.
292
+ """
293
+
294
+ pkg_logger = logging.getLogger(logger_name)
295
+ old_stdout = sys.stdout
296
+ old_stderr = sys.stderr
297
+
298
+ out_stream = _LogStreamToLogger(
299
+ logger=pkg_logger,
300
+ level=stdout_level,
301
+ carriage_return_rate_limit_s=carriage_return_rate_limit_s,
302
+ )
303
+ err_stream = _LogStreamToLogger(
304
+ logger=pkg_logger,
305
+ level=stderr_level,
306
+ carriage_return_rate_limit_s=carriage_return_rate_limit_s,
307
+ )
308
+ try:
309
+ sys.stdout = out_stream
310
+ sys.stderr = err_stream
311
+ yield
312
+ finally:
313
+ try:
314
+ out_stream.flush()
315
+ finally:
316
+ sys.stdout = old_stdout
317
+ try:
318
+ err_stream.flush()
319
+ finally:
320
+ sys.stderr = old_stderr
321
+
322
+
323
+ # Install record defaults at import time so early log emissions are safe
324
+ # even before configure_logging() is called by library entrypoints.
325
+ ensure_log_record_defaults()
metalsurfer/_utils.py ADDED
@@ -0,0 +1,13 @@
1
+ """Internal utilities shared across metalsurfer sub-packages."""
2
+
3
+ from math import isfinite
4
+
5
+
6
+ def is_finite_number(value: object) -> bool:
7
+ """Return True if *value* converts to a finite float."""
8
+ if not isinstance(value, (int, float, str)):
9
+ return False
10
+ try:
11
+ return bool(isfinite(float(value)))
12
+ except (TypeError, ValueError):
13
+ return False