phaserEM 0.1__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 (67) hide show
  1. phaser/__init__.py +0 -0
  2. phaser/__main__.py +5 -0
  3. phaser/engines/common/__init__.py +0 -0
  4. phaser/engines/common/noise_models.py +113 -0
  5. phaser/engines/common/output.py +189 -0
  6. phaser/engines/common/position_correction.py +62 -0
  7. phaser/engines/common/regularizers.py +403 -0
  8. phaser/engines/common/simulation.py +270 -0
  9. phaser/engines/conventional/__init__.py +0 -0
  10. phaser/engines/conventional/run.py +142 -0
  11. phaser/engines/conventional/solvers.py +476 -0
  12. phaser/engines/gradient/run.py +451 -0
  13. phaser/engines/gradient/solvers.py +139 -0
  14. phaser/execute.py +371 -0
  15. phaser/hooks/__init__.py +158 -0
  16. phaser/hooks/hook.py +159 -0
  17. phaser/hooks/io/empad.py +88 -0
  18. phaser/hooks/object.py +25 -0
  19. phaser/hooks/preprocessing.py +133 -0
  20. phaser/hooks/probe.py +24 -0
  21. phaser/hooks/regularization.py +97 -0
  22. phaser/hooks/scan.py +27 -0
  23. phaser/hooks/schedule.py +75 -0
  24. phaser/hooks/solver.py +169 -0
  25. phaser/io/__init__.py +0 -0
  26. phaser/io/empad.py +212 -0
  27. phaser/main.py +92 -0
  28. phaser/plan.py +184 -0
  29. phaser/py.typed +0 -0
  30. phaser/state.py +249 -0
  31. phaser/types.py +305 -0
  32. phaser/utils/__init__.py +0 -0
  33. phaser/utils/_cuda_kernels.py +213 -0
  34. phaser/utils/_jax_kernels.py +98 -0
  35. phaser/utils/analysis.py +263 -0
  36. phaser/utils/image.py +201 -0
  37. phaser/utils/io.py +402 -0
  38. phaser/utils/misc.py +295 -0
  39. phaser/utils/num.py +800 -0
  40. phaser/utils/object.py +578 -0
  41. phaser/utils/optics.py +377 -0
  42. phaser/utils/physics.py +88 -0
  43. phaser/utils/plotting.py +699 -0
  44. phaser/utils/scan.py +60 -0
  45. phaser/web/__init__.py +0 -0
  46. phaser/web/dist/03510a839ccb97b0da9f.module.wasm +0 -0
  47. phaser/web/dist/9573273f862f4f5d9644.module.wasm +0 -0
  48. phaser/web/dist/bundle-dashboard.js +712 -0
  49. phaser/web/dist/bundle-manager.js +210 -0
  50. phaser/web/dist/bundle-vendors-node_modules_wasm-array_wasm_array_js.js +106 -0
  51. phaser/web/dist/style.css +152 -0
  52. phaser/web/notebook.py +269 -0
  53. phaser/web/routes.py +211 -0
  54. phaser/web/server.py +540 -0
  55. phaser/web/slurm.py +180 -0
  56. phaser/web/templates/base.html +13 -0
  57. phaser/web/templates/dashboard.html +14 -0
  58. phaser/web/templates/manager.html +19 -0
  59. phaser/web/types.py +255 -0
  60. phaser/web/util.py +116 -0
  61. phaser/web/worker.py +200 -0
  62. phaserem-0.1.dist-info/METADATA +121 -0
  63. phaserem-0.1.dist-info/RECORD +67 -0
  64. phaserem-0.1.dist-info/WHEEL +5 -0
  65. phaserem-0.1.dist-info/entry_points.txt +2 -0
  66. phaserem-0.1.dist-info/licenses/LICENSE.txt +373 -0
  67. phaserem-0.1.dist-info/top_level.txt +1 -0
phaser/state.py ADDED
@@ -0,0 +1,249 @@
1
+ import typing as t
2
+
3
+ import numpy
4
+ from numpy.typing import NDArray
5
+ from typing_extensions import Self
6
+
7
+ from phaser.utils.num import Sampling, to_numpy, get_array_module, Float
8
+ from phaser.utils.misc import jax_dataclass
9
+ from phaser.utils.object import ObjectSampling
10
+
11
+ if t.TYPE_CHECKING:
12
+ from phaser.utils.io import HdfLike
13
+ from phaser.utils.image import _BoundaryMode
14
+
15
+
16
+ @jax_dataclass
17
+ class Patterns():
18
+ patterns: NDArray[numpy.floating]
19
+ """Raw diffraction patterns, with 0-frequency sample in corner"""
20
+ pattern_mask: NDArray[numpy.floating]
21
+ """Mask indicating which portions of the diffraction patterns contain data."""
22
+
23
+ def to_numpy(self) -> Self:
24
+ return self.__class__(
25
+ to_numpy(self.patterns), to_numpy(self.pattern_mask)
26
+ )
27
+
28
+
29
+ @jax_dataclass
30
+ class IterState():
31
+ engine_num: int
32
+ """Engine number. 1-indexed (0 means before any reconstruction)."""
33
+ engine_iter: int
34
+ """Iteration number on this engine. 1-indexed (0 means before any iterations)."""
35
+ total_iter: int
36
+ """Total iteration number. 1-indexed (0 means before any iterations)."""
37
+
38
+ def to_numpy(self) -> Self:
39
+ return self.__class__(
40
+ int(self.engine_num), int(self.engine_iter), int(self.total_iter)
41
+ )
42
+
43
+ def copy(self) -> Self:
44
+ import copy
45
+ return copy.deepcopy(self)
46
+
47
+ @staticmethod
48
+ def empty() -> 'IterState':
49
+ return IterState(0, 0, 0)
50
+
51
+
52
+ @jax_dataclass(static_fields=('sampling',))
53
+ class ProbeState():
54
+ sampling: Sampling
55
+ """Probe coordinate system. See `Sampling` for more details."""
56
+ data: NDArray[numpy.complexfloating]
57
+ """Probe wavefunction, in realspace. Shape (modes, y, x)"""
58
+
59
+ def resample(
60
+ self, new_samp: Sampling,
61
+ rotation: float = 0.0,
62
+ order: int = 1,
63
+ mode: '_BoundaryMode' = 'grid-constant',
64
+ ) -> Self:
65
+ new_data = self.sampling.resample(
66
+ self.data, new_samp,
67
+ rotation=rotation,
68
+ order=order,
69
+ mode=mode,
70
+ )
71
+ return self.__class__(new_samp, new_data)
72
+
73
+ def to_xp(self, xp: t.Any) -> Self:
74
+ return self.__class__(
75
+ self.sampling, xp.array(self.data)
76
+ )
77
+
78
+ def to_numpy(self) -> Self:
79
+ return self.__class__(
80
+ self.sampling, to_numpy(self.data)
81
+ )
82
+
83
+ def copy(self) -> Self:
84
+ import copy
85
+ return copy.deepcopy(self)
86
+
87
+
88
+ @jax_dataclass(static_fields=('sampling',))
89
+ class ObjectState():
90
+ sampling: ObjectSampling
91
+ """Object coordinate system. See `ObjectSampling` for more details."""
92
+ data: NDArray[numpy.complexfloating]
93
+ """Object wavefunction. Shape (z, y, x)"""
94
+ thicknesses: NDArray[numpy.floating]
95
+ """
96
+ Slice thicknesses (in length units).
97
+ Length < 2 for single slice, equal to the number of slices otherwise.
98
+ """
99
+
100
+ def to_xp(self, xp: t.Any) -> Self:
101
+ return self.__class__(
102
+ self.sampling, xp.array(self.data), xp.array(self.thicknesses)
103
+ )
104
+
105
+ def to_numpy(self) -> Self:
106
+ return self.__class__(
107
+ self.sampling, to_numpy(self.data), to_numpy(self.thicknesses)
108
+ )
109
+
110
+ def zs(self) -> NDArray[numpy.floating]:
111
+ xp = get_array_module(self.thicknesses)
112
+ if len(self.thicknesses) < 2:
113
+ return xp.array([0.], dtype=self.thicknesses.dtype)
114
+ return xp.cumsum(self.thicknesses) - self.thicknesses
115
+
116
+ def copy(self) -> Self:
117
+ import copy
118
+ return copy.deepcopy(self)
119
+
120
+
121
+ @jax_dataclass
122
+ class ProgressState:
123
+ iters: NDArray[numpy.integer]
124
+ """Iterations error measurements were taken at."""
125
+ detector_errors: NDArray[numpy.floating]
126
+ """Detector error measurements at those iterations"""
127
+
128
+ def to_numpy(self) -> Self:
129
+ return self.__class__(
130
+ to_numpy(self.iters), to_numpy(self.detector_errors)
131
+ )
132
+
133
+ def copy(self) -> Self:
134
+ import copy
135
+ return copy.deepcopy(self)
136
+
137
+ @staticmethod
138
+ def empty() -> 'ProgressState':
139
+ return ProgressState(
140
+ numpy.array([], dtype=numpy.uint64),
141
+ numpy.array([], dtype=numpy.float64),
142
+ )
143
+
144
+ # TODO: this is a hack to prevent JIT recompilation.
145
+ def __hash__(self) -> int:
146
+ return id(self)
147
+
148
+ def __eq__(self, other: t.Any) -> bool:
149
+ if type(self) is not type(other):
150
+ return False
151
+ xp = get_array_module(self.iters, other.iters)
152
+ return (
153
+ xp.array_equal(self.iters, other.iters) and
154
+ xp.array_equal(self.detector_errors, other.detector_errors)
155
+ )
156
+
157
+ @jax_dataclass(kw_only=True, static_fields=('progress',))
158
+ class ReconsState:
159
+ iter: IterState
160
+ wavelength: Float
161
+
162
+ probe: ProbeState
163
+ object: ObjectState
164
+ scan: NDArray[numpy.floating]
165
+ """Scan coordinates (y, x), in length units. Shape (..., 2)"""
166
+ progress: ProgressState
167
+
168
+ def to_xp(self, xp: t.Any) -> Self:
169
+ return self.__class__(
170
+ iter=self.iter,
171
+ probe=self.probe.to_xp(xp),
172
+ object=self.object.to_xp(xp),
173
+ scan=xp.array(self.scan),
174
+ progress=self.progress,
175
+ wavelength=self.wavelength,
176
+ )
177
+
178
+ def to_numpy(self) -> Self:
179
+ return self.__class__(
180
+ iter=self.iter.to_numpy(),
181
+ probe=self.probe.to_numpy(),
182
+ object=self.object.to_numpy(),
183
+ scan=to_numpy(self.scan),
184
+ progress=self.progress.to_numpy(),
185
+ wavelength=float(self.wavelength),
186
+ )
187
+
188
+ def copy(self) -> Self:
189
+ import copy
190
+ return copy.deepcopy(self)
191
+
192
+ def write_hdf5(self, file: 'HdfLike'):
193
+ from phaser.utils.io import hdf5_write_state
194
+ hdf5_write_state(self, file)
195
+
196
+ @staticmethod
197
+ def read_hdf5(file: 'HdfLike') -> 'ReconsState':
198
+ from phaser.utils.io import hdf5_read_state
199
+ return hdf5_read_state(file).to_complete()
200
+
201
+
202
+ @jax_dataclass(kw_only=True, static_fields=('progress',))
203
+ class PartialReconsState:
204
+ iter: t.Optional[IterState] = None
205
+ wavelength: t.Optional[Float] = None
206
+
207
+ probe: t.Optional[ProbeState] = None
208
+ object: t.Optional[ObjectState] = None
209
+ scan: t.Optional[NDArray[numpy.floating]] = None
210
+ """Scan coordinates (y, x), in length units. Shape (..., 2)"""
211
+ progress: t.Optional[ProgressState] = None
212
+
213
+ def to_numpy(self) -> Self:
214
+ return self.__class__(
215
+ iter=self.iter.to_numpy() if self.iter is not None else None,
216
+ probe=self.probe.to_numpy() if self.probe is not None else None,
217
+ object=self.object.to_numpy() if self.object is not None else None,
218
+ scan=to_numpy(self.scan) if self.scan is not None else None,
219
+ progress=self.progress.to_numpy() if self.progress is not None else None,
220
+ wavelength=float(self.wavelength) if self.wavelength is not None else None,
221
+ )
222
+
223
+ def to_complete(self) -> ReconsState:
224
+ missing = tuple(filter(lambda k: getattr(self, k) is None, ('probe', 'object', 'scan', 'wavelength')))
225
+ if len(missing):
226
+ raise ValueError(f"ReconsState missing {', '.join(map(repr, missing))}")
227
+
228
+ progress = self.progress if self.progress is not None else ProgressState.empty()
229
+ iter = self.iter if self.iter is not None else IterState.empty()
230
+
231
+ return ReconsState(
232
+ wavelength=t.cast(Float, self.wavelength),
233
+ probe=t.cast(ProbeState, self.probe),
234
+ object=t.cast(ObjectState, self.object),
235
+ scan=t.cast(NDArray[numpy.floating], self.scan),
236
+ progress=progress, iter=iter,
237
+ )
238
+
239
+ def write_hdf5(self, file: 'HdfLike'):
240
+ from phaser.utils.io import hdf5_write_state
241
+ hdf5_write_state(self, file)
242
+
243
+ @staticmethod
244
+ def read_hdf5(file: 'HdfLike') -> 'PartialReconsState':
245
+ from phaser.utils.io import hdf5_read_state
246
+ return hdf5_read_state(file)
247
+
248
+
249
+ StateObserver: t.TypeAlias = t.Callable[[t.Union[ReconsState, PartialReconsState]], t.Any]
phaser/types.py ADDED
@@ -0,0 +1,305 @@
1
+ from dataclasses import dataclass
2
+ from functools import lru_cache
3
+ import typing as t
4
+
5
+ import numpy
6
+ import pane
7
+ from pane.converters import Converter, make_converter, ConverterHandlers, ErrorNode
8
+ from pane.annotations import ConvertAnnotation
9
+ from pane.errors import ParseInterrupt, WrongTypeError
10
+ from pane.util import pluralize, list_phrase
11
+ from typing_extensions import Self
12
+
13
+ if t.TYPE_CHECKING:
14
+ from phaser.hooks.schedule import FlagArgs, FlagLike, ScheduleLike
15
+
16
+ T = t.TypeVar('T')
17
+
18
+ @t.overload
19
+ def cast_length(val: t.Iterable[T], n: t.Literal[5]) -> t.Tuple[T, T, T, T, T]:
20
+ ...
21
+
22
+ @t.overload
23
+ def cast_length(val: t.Iterable[T], n: t.Literal[4]) -> t.Tuple[T, T, T, T]:
24
+ ...
25
+
26
+ @t.overload
27
+ def cast_length(val: t.Iterable[T], n: t.Literal[3]) -> t.Tuple[T, T, T]:
28
+ ...
29
+
30
+ @t.overload
31
+ def cast_length(val: t.Iterable[T], n: t.Literal[2]) -> t.Tuple[T, T]:
32
+ ...
33
+
34
+ @t.overload
35
+ def cast_length(val: t.Iterable[T], n: t.Literal[1]) -> t.Tuple[T]:
36
+ ...
37
+
38
+ def cast_length(val: t.Iterable[T], n: int) -> t.Tuple[T, ...]:
39
+ return tuple(val)
40
+
41
+
42
+ class _EmptyDictAnnotation(ConvertAnnotation, Converter[t.Dict[t.NoReturn, t.NoReturn]]):
43
+ def _converter(self, inner_type: t.Any, *, handlers: ConverterHandlers):
44
+ return self
45
+
46
+ def __eq__(self, other):
47
+ return self.__class__ is other.__class__
48
+
49
+ def __hash__(self) -> int:
50
+ return hash(self.__class__.__name__)
51
+
52
+ def expected(self, plural: bool = False) -> str:
53
+ return pluralize("empty dict", plural, article='an')
54
+
55
+ def try_convert(self, val: t.Any) -> t.Dict[t.NoReturn, t.NoReturn]:
56
+ if isinstance(val, (dict, t.Mapping)) and len(val) == 0:
57
+ return {}
58
+ raise ParseInterrupt()
59
+
60
+ def collect_errors(self, val: t.Any) -> t.Optional[WrongTypeError]:
61
+ if isinstance(val, dict) and len(val) == 0:
62
+ return None
63
+ return WrongTypeError(self.expected(), val)
64
+
65
+
66
+ class _ReconsVarsAnnotation(ConvertAnnotation):
67
+ def _converter(self, inner_type: t.Any, *, handlers: ConverterHandlers):
68
+ return _ReconsVarsConverter(inner_type, handlers)
69
+
70
+ def __eq__(self, other):
71
+ return self.__class__ is other.__class__
72
+
73
+ def __hash__(self) -> int:
74
+ return hash(self.__class__.__name__)
75
+
76
+
77
+ BackendName: t.TypeAlias = t.Literal['cuda', 'cupy', 'jax', 'cpu', 'numpy']
78
+ ReconsVar: t.TypeAlias = t.Literal['object', 'probe', 'positions']
79
+
80
+ ReconsVars: t.TypeAlias = t.Annotated[t.FrozenSet[ReconsVar], _ReconsVarsAnnotation()]
81
+ EmptyDict: t.TypeAlias = t.Annotated[t.Dict[t.NoReturn, t.NoReturn], _EmptyDictAnnotation()]
82
+
83
+
84
+ class Cancelled(BaseException):
85
+ ...
86
+
87
+
88
+ class Dataclass(pane.PaneBase, kw_only=True, allow_extra=False):
89
+ ...
90
+
91
+
92
+ class SliceList(Dataclass):
93
+ thicknesses: t.List[float]
94
+
95
+ class SliceStep(Dataclass):
96
+ n: int
97
+ slice_thickness: float
98
+
99
+ @property
100
+ def zs(self) -> t.List[float]:
101
+ return [float(z) for z in numpy.arange(self.n) * self.slice_thickness]
102
+
103
+ @property
104
+ def thicknesses(self) -> t.List[float]:
105
+ return [self.slice_thickness] * self.n
106
+
107
+ class SliceTotal(Dataclass):
108
+ n: int
109
+ total_thickness: float
110
+
111
+ @property
112
+ def zs(self) -> t.List[float]:
113
+ return [float(z) for z in numpy.arange(self.n) * self.total_thickness / self.n]
114
+
115
+ @property
116
+ def thicknesses(self) -> t.List[float]:
117
+ return [self.total_thickness / self.n] * self.n
118
+
119
+ Slices: t.TypeAlias = t.Union[SliceList, SliceStep, SliceTotal]
120
+
121
+
122
+ class Flag(Dataclass):
123
+ after: int = 0
124
+ every: int = 1
125
+ before: t.Optional[int] = None
126
+
127
+ def any_true(self, niter: int) -> bool:
128
+ return (
129
+ self.after < niter
130
+ and (self.before is None or self.before > 0)
131
+ )
132
+
133
+ def __call__(self, args: 'FlagArgs') -> bool:
134
+ i = args['state'].iter.engine_iter
135
+ return (
136
+ (self.before is None or i < self.before)
137
+ and i > self.after
138
+ and (i - self.after) % self.every == 0
139
+ )
140
+
141
+ def resolve(self) -> Self:
142
+ return self
143
+
144
+
145
+ class _ConstFlag:
146
+ def __init__(self, val: bool):
147
+ self.val = val
148
+
149
+ def __call__(self, args: 'FlagArgs') -> bool:
150
+ return self.val
151
+
152
+
153
+ @lru_cache
154
+ def process_flag(flag: 'FlagLike') -> t.Callable[['FlagArgs'], bool]:
155
+ if isinstance(flag, bool):
156
+ return _ConstFlag(flag)
157
+ return flag
158
+
159
+
160
+ @lru_cache
161
+ def process_schedule(schedule: 'ScheduleLike') -> t.Callable[['FlagArgs'], float]:
162
+ if isinstance(schedule, (int, float)):
163
+ return lambda _: schedule
164
+ return schedule
165
+
166
+
167
+ def flag_any_true(flag: t.Callable[['FlagArgs'], bool], niter: int) -> bool:
168
+ if isinstance(flag, Flag):
169
+ return flag.any_true(niter)
170
+ elif isinstance(flag, _ConstFlag):
171
+ return flag.val
172
+ # assume flag will return true
173
+ return True
174
+
175
+
176
+ @dataclass(init=False, frozen=True)
177
+ class IsVersion(ConvertAnnotation):
178
+ min: t.Optional[t.Tuple[int, ...]] = None
179
+ max: t.Optional[t.Tuple[int, ...]] = None
180
+
181
+ def __init__(self, *,
182
+ min: t.Union[str, t.Tuple[int, ...], None] = None,
183
+ max: t.Union[str, t.Tuple[int, ...], None] = None,
184
+ exactly: t.Union[str, t.Tuple[int, ...], None] = None):
185
+ if exactly is not None:
186
+ if min is not None or max is not None:
187
+ raise TypeError("'exactly' cannot be specified with 'min' or 'max'")
188
+ min = max = _VersionConverter.parse_version(exactly)
189
+ if min is not None:
190
+ min = _VersionConverter.parse_version(min)
191
+ if max is not None:
192
+ max = _VersionConverter.parse_version(max)
193
+ object.__setattr__(self, 'min', min)
194
+ object.__setattr__(self, 'max', max)
195
+
196
+ def _converter(self, inner_type: t.Any, *,
197
+ handlers: ConverterHandlers):
198
+ return _VersionConverter(inner_type, handlers, min=self.min, max=self.max)
199
+
200
+
201
+ Version: t.TypeAlias = t.Annotated[str, IsVersion()]
202
+
203
+
204
+ class _VersionConverter(Converter[t.Any]):
205
+ def __init__(self, inner_type: t.Any, handlers: ConverterHandlers,
206
+ min: t.Optional[t.Tuple[int, ...]] = None,
207
+ max: t.Optional[t.Tuple[int, ...]] = None):
208
+ self.inner = make_converter(inner_type, handlers)
209
+ self.min = min
210
+ self.max = max
211
+
212
+ def expected(self, plural: bool = False) -> str:
213
+ if self.min is not None and self.max is not None:
214
+ min_s = '.'.join(map(str, self.min))
215
+ max_s = '.'.join(map(str, self.max))
216
+ if self.min == self.max:
217
+ return f"{pluralize('version', plural)} {min_s}"
218
+ return f"{pluralize('version', plural)} between {min_s} and {max_s}"
219
+ elif self.min is not None:
220
+ min_s = '.'.join(map(str, self.min))
221
+ return f"{pluralize('version', plural)} at least {min_s}"
222
+ elif self.max is not None:
223
+ max_s = '.'.join(map(str, self.max))
224
+ return f"{pluralize('version', plural)} at most {max_s}"
225
+ else:
226
+ return f"version {pluralize('string', plural)}"
227
+
228
+ def into_data(self, val: t.Any) -> str:
229
+ return str(val)
230
+
231
+ @staticmethod
232
+ def parse_version(val: t.Union[str, t.Tuple[int, ...]]) -> t.Tuple[int, ...]:
233
+ def to_int(seg: t.Union[int, str]) -> int:
234
+ if isinstance(seg, int):
235
+ return seg
236
+ seg = seg.strip()
237
+ if not seg.isdigit():
238
+ raise ValueError()
239
+ return int(seg)
240
+
241
+ return tuple(map(to_int, val.split('.') if isinstance(val, str) else val))
242
+
243
+ def check_version(self, val: t.Tuple[int, ...]):
244
+ if self.min == self.max:
245
+ if self.min is not None and val != self.min:
246
+ raise ValueError(f"Version {'.'.join(map(str, val))} is not supported version {'.'.join(map(str, self.min))}")
247
+ elif self.min is not None and val < self.min:
248
+ raise ValueError(f"Version {'.'.join(map(str, val))} less than minimum supported version {'.'.join(map(str, self.min))}")
249
+ elif self.max is not None and val > self.max:
250
+ raise ValueError(f"Version {'.'.join(map(str, val))} greater than maximum supported version {'.'.join(map(str, self.max))}")
251
+
252
+ def try_convert(self, val: t.Any) -> str:
253
+ s = self.inner.try_convert(val)
254
+ try:
255
+ version = self.parse_version(val)
256
+ self.check_version(version)
257
+ except ValueError:
258
+ raise ParseInterrupt()
259
+ return s
260
+
261
+ def collect_errors(self, val: t.Any) -> t.Optional[ErrorNode]:
262
+ try:
263
+ s = self.inner.try_convert(val)
264
+ except ParseInterrupt:
265
+ return WrongTypeError(self.expected(), val)
266
+ try:
267
+ version = self.parse_version(s)
268
+ except ValueError:
269
+ return WrongTypeError(self.expected(), val, info='Invalid version string')
270
+ try:
271
+ self.check_version(version)
272
+ except ValueError as e:
273
+ return WrongTypeError(self.expected(), val, info=e.args[0])
274
+ return None
275
+
276
+
277
+ class _ReconsVarsConverter(Converter[t.FrozenSet[ReconsVar]]):
278
+ def __init__(self, ty: type, handlers: ConverterHandlers):
279
+ self.inner = make_converter(ty, handlers)
280
+
281
+ def expected(self, plural: bool = False) -> str:
282
+ known_params = t.get_args(ReconsVar)
283
+ return f"{pluralize('set', plural, article='a')} of comma-separated " \
284
+ f"variables ({list_phrase(tuple(map(repr, known_params)))})"
285
+
286
+ def into_data(self, val: t.Any) -> str:
287
+ return ", ".join(self.inner.into_data(val)) # type: ignore
288
+
289
+ def try_convert(self, val: t.Any) -> t.FrozenSet[ReconsVar]:
290
+ if isinstance(val, str):
291
+ val = tuple(v.strip() for v in val.split(","))
292
+
293
+ return self.inner.try_convert(val)
294
+
295
+ def collect_errors(self, val: t.Any) -> t.Optional[ErrorNode]:
296
+ if isinstance(val, str):
297
+ val = tuple(v.strip() for v in val.split(","))
298
+
299
+ return self.inner.collect_errors(val)
300
+
301
+
302
+ __all__ = [
303
+ 'BackendName', 'Dataclass', 'Slices', 'Flag',
304
+ 'process_flag', 'flag_any_true',
305
+ ]
File without changes