netloader 3.11.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.
netloader/__init__.py ADDED
@@ -0,0 +1,61 @@
1
+ """
2
+ Package information, creates the logger and adds netloader classes to PyTorch safe globals
3
+ """
4
+ import sys
5
+ import logging
6
+ import warnings
7
+
8
+
9
+ __version__ = '3.11.0'
10
+ __author__ = 'Ethan Tregidga'
11
+ logging.basicConfig(format='%(levelname)s: %(message)s', level=logging.WARNING)
12
+ warnings.filterwarnings('once', category=DeprecationWarning, module=r'^netloader(\.|$)')
13
+ warnings.filterwarnings(
14
+ 'once',
15
+ category=PendingDeprecationWarning,
16
+ module=r'^netloader(\.|$)',
17
+ )
18
+
19
+
20
+ if sys.version_info < (3, 12):
21
+ warnings.warn(
22
+ f'Python version ({sys.version_info.major}.{sys.version_info.minor}) is '
23
+ f'deprecated, please upgrade to Python 3.12 or higher',
24
+ DeprecationWarning,
25
+ )
26
+
27
+ try:
28
+ import torch
29
+
30
+ from netloader.utils import safe_globals
31
+ from netloader import (
32
+ layers,
33
+ models,
34
+ architectures,
35
+ utils,
36
+ data,
37
+ loss_funcs,
38
+ network,
39
+ schedulers,
40
+ transforms,
41
+ )
42
+
43
+
44
+ # Adds PyTorch Network Loader classes to list of safe PyTorch classes when loading saved
45
+ # architectures
46
+ safe_globals(__name__, [models, architectures, loss_funcs, schedulers, transforms])
47
+ torch.serialization.add_safe_globals([network.Network, network.CompatibleNetwork])
48
+
49
+ __all__ = [
50
+ 'utils',
51
+ 'layers',
52
+ 'models',
53
+ 'architectures',
54
+ 'data',
55
+ 'network',
56
+ 'loss_funcs',
57
+ 'schedulers',
58
+ 'transforms',
59
+ ]
60
+ except (ModuleNotFoundError, ImportError):
61
+ pass
@@ -0,0 +1,48 @@
1
+ """
2
+ Collects all architectures
3
+ """
4
+ from typing import Any
5
+ from warnings import warn
6
+
7
+ from netloader.architectures.utils import UtilityMixin
8
+ from netloader.architectures.base import BaseArchitecture, load_arch
9
+ from netloader.architectures.encoder_decoder import Autoencoder, Decoder, Encoder
10
+
11
+
12
+ __all__ = [
13
+ 'BaseArchitecture',
14
+ 'UtilityMixin',
15
+ 'Autoencoder',
16
+ 'Decoder',
17
+ 'Encoder',
18
+ 'load_arch',
19
+ ]
20
+
21
+ _optional_imports: dict[str, set[str]] = {'flow': {'NormFlow', 'NormFlowEncoder'}}
22
+
23
+ try:
24
+ from netloader.architectures.flows import NormFlow, NormFlowEncoder
25
+ __all__.extend(_optional_imports['flow'])
26
+ except (ModuleNotFoundError, ImportError) as e:
27
+ pass
28
+
29
+
30
+ def __getattr__(name: str) -> Any:
31
+ if name == 'BaseNetwork':
32
+ warn(
33
+ 'BaseNetwork is deprecated, please use BaseArchitecture instead',
34
+ DeprecationWarning,
35
+ stacklevel=2,
36
+ )
37
+ return BaseArchitecture
38
+ if name == 'load_net':
39
+ warn(
40
+ 'load_net is deprecated, please use load_arch instead',
41
+ DeprecationWarning,
42
+ stacklevel=2,
43
+ )
44
+ return load_arch
45
+ if name not in __all__ and name in _optional_imports['flow']:
46
+ raise ImportError(f'Cannot import {name}, normalising flow architectures require the '
47
+ f'package Zuko')
48
+ raise AttributeError(f"module {__name__} has no attribute {name}")