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.
- metalsurfer/__init__.py +147 -0
- metalsurfer/_logging.py +325 -0
- metalsurfer/_utils.py +13 -0
- metalsurfer/campaigns.py +583 -0
- metalsurfer/config.py +695 -0
- metalsurfer/conformers.py +249 -0
- metalsurfer/exceptions.py +29 -0
- metalsurfer/filters.py +604 -0
- metalsurfer/io_results.py +751 -0
- metalsurfer/ml/__init__.py +11 -0
- metalsurfer/ml/bayesian.py +927 -0
- metalsurfer/ml/dataset.py +217 -0
- metalsurfer/ml/features.py +147 -0
- metalsurfer/ml/predict.py +123 -0
- metalsurfer/ml/regression.py +307 -0
- metalsurfer/ml/reproduce.py +158 -0
- metalsurfer/ml/schema.py +579 -0
- metalsurfer/models.py +773 -0
- metalsurfer/optimization.py +1300 -0
- metalsurfer/placement/__init__.py +64 -0
- metalsurfer/placement/_constants.py +211 -0
- metalsurfer/placement/_material.py +59 -0
- metalsurfer/placement/generators.py +1435 -0
- metalsurfer/placement/geometry.py +856 -0
- metalsurfer/placement/policy.py +206 -0
- metalsurfer/placement/sites.py +1815 -0
- metalsurfer/py.typed +0 -0
- metalsurfer/surface_prep/__init__.py +69 -0
- metalsurfer/surface_prep/prep.py +372 -0
- metalsurfer/surfaces.py +1004 -0
- metalsurfer/symmetry.py +460 -0
- metalsurfer/workflow/__init__.py +15 -0
- metalsurfer/workflow/bayesian.py +601 -0
- metalsurfer/workflow/core.py +498 -0
- metalsurfer/workflow/reference.py +92 -0
- metalsurfer/workflow/saturation.py +913 -0
- metalsurfer/workflow/shared.py +1070 -0
- metalsurfer-0.3.0.dist-info/METADATA +41 -0
- metalsurfer-0.3.0.dist-info/RECORD +41 -0
- metalsurfer-0.3.0.dist-info/WHEEL +5 -0
- metalsurfer-0.3.0.dist-info/top_level.txt +1 -0
metalsurfer/__init__.py
ADDED
|
@@ -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}")
|
metalsurfer/_logging.py
ADDED
|
@@ -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
|