Fola-lad commited on
Commit
7057729
·
1 Parent(s): 9324a92

wire up FFN model — model_def, label map fix, TF requirement

Browse files
Files changed (4) hide show
  1. .gitignore +4 -0
  2. requirements.txt +1 -0
  3. src/model_def.py +65 -0
  4. src/streamlit_app.py +18 -9
.gitignore ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ model.keras
2
+ *.keras
3
+ __pycache__/
4
+ .DS_Store
requirements.txt CHANGED
@@ -2,3 +2,4 @@ streamlit
2
  pandas
3
  numpy
4
  scipy
 
 
2
  pandas
3
  numpy
4
  scipy
5
+ tensorflow-cpu==2.16.2
src/model_def.py ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """FeedForwardNetwork definition required for deserializing model.keras.
2
+
3
+ This must be imported before tf.keras.models.load_model() is called so
4
+ that keras can resolve the registered custom class.
5
+ """
6
+
7
+ import keras
8
+ import tensorflow as tf
9
+
10
+
11
+ @keras.saving.register_keras_serializable()
12
+ class FeedForwardNetwork(tf.keras.Model):
13
+ """Fully-connected feedforward network for 6-class HAR classification.
14
+
15
+ Architecture: Dense(512) → BN → ReLU → Dropout
16
+ Dense(256) → BN → ReLU → Dropout
17
+ Dense(128) → BN → ReLU → Dropout
18
+ Dense(6, softmax)
19
+ """
20
+
21
+ def __init__(
22
+ self,
23
+ num_features,
24
+ num_classes,
25
+ hidden_units=(512, 256, 128),
26
+ dropout_rate=0.3,
27
+ **kwargs,
28
+ ):
29
+ super().__init__(**kwargs)
30
+ self._num_features = num_features
31
+ self._num_classes = num_classes
32
+ self._hidden_units = tuple(hidden_units)
33
+ self._dropout_rate = dropout_rate
34
+
35
+ self.hidden_blocks = []
36
+ for units in hidden_units:
37
+ self.hidden_blocks.append([
38
+ tf.keras.layers.Dense(units, use_bias=False),
39
+ tf.keras.layers.BatchNormalization(),
40
+ tf.keras.layers.ReLU(),
41
+ tf.keras.layers.Dropout(dropout_rate),
42
+ ])
43
+
44
+ self.output_layer = tf.keras.layers.Dense(num_classes, activation="softmax")
45
+
46
+ def call(self, inputs, training=False):
47
+ x = inputs
48
+ for block in self.hidden_blocks:
49
+ for layer in block:
50
+ if isinstance(layer, (tf.keras.layers.BatchNormalization,
51
+ tf.keras.layers.Dropout)):
52
+ x = layer(x, training=training)
53
+ else:
54
+ x = layer(x)
55
+ return self.output_layer(x)
56
+
57
+ def get_config(self):
58
+ config = super().get_config()
59
+ config.update({
60
+ "num_features": self._num_features,
61
+ "num_classes": self._num_classes,
62
+ "hidden_units": self._hidden_units,
63
+ "dropout_rate": self._dropout_rate,
64
+ })
65
+ return config
src/streamlit_app.py CHANGED
@@ -5,12 +5,12 @@ import pandas as pd
5
  # ── Constants ──────────────────────────────────────────────────────────────
6
 
7
  LABEL_MAP = {
8
- 0: "LAYING",
9
- 1: "SITTING",
10
- 2: "STANDING",
11
- 3: "WALKING",
12
- 4: "WALKING_DOWNSTAIRS",
13
- 5: "WALKING_UPSTAIRS",
14
  }
15
 
16
  ACTIVITY_ICONS = {
@@ -35,7 +35,16 @@ EXPLANATIONS = {
35
 
36
  @st.cache_resource
37
  def load_model():
38
- return None, "no_model"
 
 
 
 
 
 
 
 
 
39
 
40
  # ── Page config ─────────────────────────────────────────────────────────────
41
 
@@ -66,8 +75,8 @@ with st.sidebar:
66
  """)
67
  st.markdown("---")
68
  st.markdown("**Model performance on test set**")
69
- st.metric("Accuracy", "")
70
- st.metric("Macro F1", "")
71
  st.markdown("---")
72
  st.caption("DAT606 Group Assignment · Pan-Atlantic University")
73
 
 
5
  # ── Constants ──────────────────────────────────────────────────────────────
6
 
7
  LABEL_MAP = {
8
+ 0: "WALKING",
9
+ 1: "WALKING_UPSTAIRS",
10
+ 2: "WALKING_DOWNSTAIRS",
11
+ 3: "SITTING",
12
+ 4: "STANDING",
13
+ 5: "LAYING",
14
  }
15
 
16
  ACTIVITY_ICONS = {
 
35
 
36
  @st.cache_resource
37
  def load_model():
38
+ try:
39
+ import tensorflow as tf
40
+ from model_def import FeedForwardNetwork
41
+ model = tf.keras.models.load_model(
42
+ "model.keras",
43
+ custom_objects={"FeedForwardNetwork": FeedForwardNetwork},
44
+ )
45
+ return model, "ready"
46
+ except Exception:
47
+ return None, "no_model"
48
 
49
  # ── Page config ─────────────────────────────────────────────────────────────
50
 
 
75
  """)
76
  st.markdown("---")
77
  st.markdown("**Model performance on test set**")
78
+ st.metric("Architecture", "FFN 512→256→128")
79
+ st.metric("Status", "FFN live · CNN pending")
80
  st.markdown("---")
81
  st.caption("DAT606 Group Assignment · Pan-Atlantic University")
82