popboat1 commited on
Commit
12cd8e3
·
1 Parent(s): 98701be

Implement SafeDense interceptor to bypass Keras 3 h5 loading bug

Browse files
Files changed (2) hide show
  1. requirements.txt +1 -1
  2. src/api/api.py +17 -2
requirements.txt CHANGED
@@ -4,4 +4,4 @@ python-multipart
4
  opencv-python
5
  numpy
6
  matplotlib
7
- tensorflow-cpu==2.15.0
 
4
  opencv-python
5
  numpy
6
  matplotlib
7
+ tensorflow-cpu
src/api/api.py CHANGED
@@ -29,11 +29,26 @@ app.add_middleware(
29
 
30
  os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'
31
 
 
 
 
 
 
 
 
 
 
 
32
  BASE_DIR = os.path.dirname(os.path.abspath(__file__))
33
  MODEL_PATH = os.path.join(BASE_DIR, 'alexnet_cifar10_keras.h5')
34
 
35
- print("Loading AlexNet from scratch...")
36
- model = tf.keras.models.load_model(MODEL_PATH)
 
 
 
 
 
37
 
38
  conv_layers = [layer for layer in model.layers if isinstance(layer, tf.keras.layers.Conv2D)]
39
  feature_extractor = tf.keras.Model(inputs=model.inputs, outputs=[layer.output for layer in conv_layers])
 
29
 
30
  os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'
31
 
32
+ class SafeDense(tf.keras.layers.Dense):
33
+ def __init__(self, **kwargs):
34
+ kwargs.pop('quantization_config', None)
35
+ super().__init__(**kwargs)
36
+
37
+ class SafeConv2D(tf.keras.layers.Conv2D):
38
+ def __init__(self, **kwargs):
39
+ kwargs.pop('quantization_config', None)
40
+ super().__init__(**kwargs)
41
+
42
  BASE_DIR = os.path.dirname(os.path.abspath(__file__))
43
  MODEL_PATH = os.path.join(BASE_DIR, 'alexnet_cifar10_keras.h5')
44
 
45
+ model = tf.keras.models.load_model(
46
+ MODEL_PATH,
47
+ custom_objects={
48
+ 'Dense': SafeDense,
49
+ 'Conv2D': SafeConv2D
50
+ }
51
+ )
52
 
53
  conv_layers = [layer for layer in model.layers if isinstance(layer, tf.keras.layers.Conv2D)]
54
  feature_extractor = tf.keras.Model(inputs=model.inputs, outputs=[layer.output for layer in conv_layers])