Recursively convert configuration values into YAML-friendly objects.
(
value: Any,
*,
custom_serializers: Optional[Iterable[Serializer]] = None,
)
| 212 | |
| 213 | |
| 214 | def serialize_config_value( |
| 215 | value: Any, |
| 216 | *, |
| 217 | custom_serializers: Optional[Iterable[Serializer]] = None, |
| 218 | ) -> Any: |
| 219 | """Recursively convert configuration values into YAML-friendly objects.""" |
| 220 | |
| 221 | # 添加内置序列化器 |
| 222 | def path_serializer(obj: Any, _serialize: Callable) -> Any: |
| 223 | from pathlib import Path |
| 224 | if isinstance(obj, Path): |
| 225 | return str(obj) |
| 226 | return None |
| 227 | |
| 228 | def torch_dtype_serializer(obj: Any, _serialize: Callable) -> Any: |
| 229 | try: |
| 230 | import torch |
| 231 | if isinstance(obj, torch.dtype): |
| 232 | return str(obj) |
| 233 | except ImportError: |
| 234 | pass |
| 235 | return None |
| 236 | |
| 237 | # 将内置序列化器添加到自定义序列化器列表前面 |
| 238 | serializers = [ |
| 239 | path_serializer, |
| 240 | torch_dtype_serializer, |
| 241 | *(custom_serializers or []) |
| 242 | ] |
| 243 | |
| 244 | def _serialize(obj: Any) -> Any: |
| 245 | for serializer in serializers: |
| 246 | result = serializer(obj, _serialize) |
| 247 | if result is not None: |
| 248 | return result |
| 249 | |
| 250 | if is_dataclass(obj): |
| 251 | return {f.name: _serialize(getattr(obj, f.name)) for f in fields(obj)} |
| 252 | if isinstance(obj, Mapping): |
| 253 | return {k: _serialize(v) for k, v in obj.items()} |
| 254 | if isinstance(obj, (list, tuple)): |
| 255 | return [_serialize(v) for v in obj] |
| 256 | return obj |
| 257 | |
| 258 | return _serialize(value) |
| 259 | |
| 260 | |
| 261 | def save_config_snapshot( |
no test coverage detected