Aliazimi00 commited on
Commit
338466c
·
verified ·
1 Parent(s): 85ffab7

Update core/plot.py

Browse files
Files changed (1) hide show
  1. core/plot.py +124 -13
core/plot.py CHANGED
@@ -2,6 +2,7 @@ import matplotlib.pyplot as plt
2
  import seaborn as sns
3
  import pandas as pd
4
  import numpy as np
 
5
 
6
 
7
  def plot_forecast(result):
@@ -60,24 +61,21 @@ def plot_future_forecast(df, result, future_df):
60
 
61
  def plot_metrics_precision(result):
62
  fig, ax = plt.subplots(figsize=(8, 5))
63
- metrics = {k: v for k, v in result['metrics'].items() if k in ['R2', 'Explained Variance', 'MDA (%)'] and v is not None}
64
  if not metrics:
65
  ax.text(0.5, 0.5, "No valid precision metrics available", ha='center', va='center')
66
  ax.set_title("Precision Metrics (Model Accuracy)")
67
  return fig
68
 
69
  sns.barplot(x=list(metrics.keys()), y=list(metrics.values()), ax=ax, palette="Blues_d")
70
- # Add labels on bars
71
  for i, v in enumerate(metrics.values()):
72
- ax.text(i, v + 0.01 * max(metrics.values(), default=1), f"{v:.4f}", ha='center', va='bottom', fontsize=10)
73
 
74
- # Dynamic y-axis scaling
75
- max_val = max(metrics.values(), default=1)
76
  min_val = min(metrics.values(), default=0)
77
- ax.set_ylim(min(min_val - 0.1 * abs(min_val), -0.1), max_val + 0.2 * max_val)
78
-
79
  ax.set_title("Precision Metrics (Model Accuracy)")
80
- ax.set_ylabel("Value")
81
  ax.grid(True, axis='y')
82
  plt.tight_layout()
83
  return fig
@@ -92,16 +90,13 @@ def plot_metrics_risk(result):
92
  return fig
93
 
94
  sns.barplot(x=list(metrics.keys()), y=list(metrics.values()), ax=ax, palette="Reds_d")
95
- # Add labels on bars
96
  for i, v in enumerate(metrics.values()):
97
- ax.text(i, v + 0.01 * max(metrics.values(), default=1), f"{v:.4f}", ha='center', va='bottom', fontsize=10)
98
 
99
- # Dynamic y-axis scaling
100
  max_val = max(metrics.values(), default=1)
101
  ax.set_ylim(0, max_val + 0.2 * max_val)
102
-
103
- ax.set_title("Risk Metrics (Error Magnitude)")
104
  ax.set_ylabel("Value")
 
105
  ax.grid(True, axis='y')
106
  plt.tight_layout()
107
  return fig
@@ -119,5 +114,121 @@ def plot_loss_curve(result):
119
  ax.set_ylabel('Loss (MSE)')
120
  ax.legend()
121
  ax.grid(True)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
122
  plt.tight_layout()
123
  return fig
 
2
  import seaborn as sns
3
  import pandas as pd
4
  import numpy as np
5
+ import networkx as nx
6
 
7
 
8
  def plot_forecast(result):
 
61
 
62
  def plot_metrics_precision(result):
63
  fig, ax = plt.subplots(figsize=(8, 5))
64
+ metrics = {k: v for k, v in result['metrics'].items() if k in ['R² (%)', 'Explained Variance (%)', 'MDA (%)'] and v is not None}
65
  if not metrics:
66
  ax.text(0.5, 0.5, "No valid precision metrics available", ha='center', va='center')
67
  ax.set_title("Precision Metrics (Model Accuracy)")
68
  return fig
69
 
70
  sns.barplot(x=list(metrics.keys()), y=list(metrics.values()), ax=ax, palette="Blues_d")
 
71
  for i, v in enumerate(metrics.values()):
72
+ ax.text(i, v + 0.01 * max(metrics.values(), default=100), f"{v:.2f}%", ha='center', va='bottom', fontsize=10)
73
 
74
+ max_val = max(metrics.values(), default=100)
 
75
  min_val = min(metrics.values(), default=0)
76
+ ax.set_ylim(min(min_val - 5, -10), max_val + 10)
77
+ ax.set_ylabel("Value (%)")
78
  ax.set_title("Precision Metrics (Model Accuracy)")
 
79
  ax.grid(True, axis='y')
80
  plt.tight_layout()
81
  return fig
 
90
  return fig
91
 
92
  sns.barplot(x=list(metrics.keys()), y=list(metrics.values()), ax=ax, palette="Reds_d")
 
93
  for i, v in enumerate(metrics.values()):
94
+ ax.text(i, v + 0.01 * max(metrics.values(), default=1), f"{v:.2f}", ha='center', va='bottom', fontsize=10)
95
 
 
96
  max_val = max(metrics.values(), default=1)
97
  ax.set_ylim(0, max_val + 0.2 * max_val)
 
 
98
  ax.set_ylabel("Value")
99
+ ax.set_title("Risk Metrics (Error Magnitude)")
100
  ax.grid(True, axis='y')
101
  plt.tight_layout()
102
  return fig
 
114
  ax.set_ylabel('Loss (MSE)')
115
  ax.legend()
116
  ax.grid(True)
117
+ plt.tight_layout()
118
+ return fig
119
+
120
+
121
+ def plot_model_architecture(result):
122
+ fig, ax = plt.subplots(figsize=(10, 6))
123
+ ax.axis('off') # Hide axes for clean visualization
124
+ G = nx.DiGraph()
125
+
126
+ if "architecture" not in result:
127
+ ax.text(0.5, 0.5, "No architecture details available", ha='center', va='center', fontsize=12)
128
+ ax.set_title("Model Architecture")
129
+ plt.tight_layout()
130
+ return fig
131
+
132
+ arch = result["architecture"]
133
+ model_name = arch["model_name"]
134
+ num_layers = arch["num_layers"]
135
+ hidden_units = arch["hidden_units"]
136
+ dropout = arch["dropout"]
137
+ batch_size = arch["batch_size"]
138
+ input_size = arch["input_size"]
139
+ output_size = arch["output_size"]
140
+
141
+ # Simplify node counts for visualization (limit to 5 nodes per layer to avoid clutter)
142
+ max_nodes_display = 5
143
+ input_nodes = min(input_size, max_nodes_display)
144
+ hidden_nodes = min(hidden_units, max_nodes_display)
145
+ output_nodes = min(output_size, max_nodes_display)
146
+
147
+ # Add nodes for each layer
148
+ nodes = []
149
+ pos = {}
150
+ layer_width = 1.0 / (num_layers + 2) # Space layers evenly
151
+ y_pos = 0.5 # Center vertically
152
+
153
+ # Input layer
154
+ for i in range(input_nodes):
155
+ node = f"input_{i}"
156
+ G.add_node(node, layer="input")
157
+ pos[node] = (0, y_pos + (i - input_nodes / 2) * 0.1)
158
+ nodes.append([f"input_{i}" for i in range(input_nodes)])
159
+
160
+ # Hidden layers
161
+ for layer in range(num_layers):
162
+ layer_nodes = []
163
+ for i in range(hidden_nodes):
164
+ node = f"hidden_{layer}_{i}"
165
+ G.add_node(node, layer=f"hidden_{layer+1}")
166
+ pos[node] = ((layer + 1) * layer_width, y_pos + (i - hidden_nodes / 2) * 0.1)
167
+ layer_nodes.append(node)
168
+ nodes.append(layer_nodes)
169
+
170
+ # Output layer
171
+ output_layer_nodes = []
172
+ for i in range(output_nodes):
173
+ node = f"output_{i}"
174
+ G.add_node(node, layer="output")
175
+ pos[node] = ((num_layers + 1) * layer_width, y_pos + (i - output_nodes / 2) * 0.1)
176
+ output_layer_nodes.append(node)
177
+ nodes.append(output_layer_nodes)
178
+
179
+ # Add edges between layers
180
+ for layer in range(len(nodes) - 1):
181
+ for src in nodes[layer]:
182
+ for dst in nodes[layer + 1]:
183
+ G.add_edge(src, dst)
184
+
185
+ # Handle special cases for complex models
186
+ if model_name in ["CNNModel", "HybridModel", "CNN_GRU"]:
187
+ # Simplified representation for CNN-based models
188
+ G = nx.DiGraph()
189
+ pos = {}
190
+ nodes = []
191
+ x_pos = 0
192
+ if model_name == "CNNModel":
193
+ components = ["Input", "Conv1D (32 filters)", "Output"]
194
+ component_sizes = [input_size, 32, output_size]
195
+ elif model_name == "HybridModel":
196
+ components = ["Input", "Conv1D (32 filters)", f"BiLSTM ({num_layers} layers)", "Output"]
197
+ component_sizes = [input_size, 32, hidden_units * 2, output_size]
198
+ elif model_name == "CNN_GRU":
199
+ components = ["Input", "Conv1D (32 filters)", f"GRU ({num_layers} layers)", "Output"]
200
+ component_sizes = [input_size, 32, hidden_units, output_size]
201
+
202
+ for i, comp in enumerate(components):
203
+ G.add_node(comp, layer=comp)
204
+ pos[comp] = (i * layer_width * (num_layers + 2) / len(components), y_pos)
205
+ nodes.append([comp])
206
+ if i > 0:
207
+ G.add_edge(components[i-1], comp)
208
+
209
+ # Draw the graph
210
+ nx.draw(G, pos, ax=ax, with_labels=False, node_color='lightblue', edge_color='gray', node_size=500, arrowsize=10)
211
+
212
+ # Add layer labels
213
+ for node in G.nodes(data=True):
214
+ layer = node[1]['layer']
215
+ x, y = pos[node[0]]
216
+ if layer.startswith("hidden"):
217
+ label = f"Layer {layer.split('_')[1]}: {hidden_units} units"
218
+ elif layer == "input":
219
+ label = f"Input: {input_size} units"
220
+ elif layer == "output":
221
+ label = f"Output: {output_size} units"
222
+ else:
223
+ label = layer # For CNN/Hybrid/CNN-GRU components
224
+ ax.text(x, y + 0.05, label, ha='center', va='bottom', fontsize=8)
225
+
226
+ # Add model details as title and annotation
227
+ title = f"{model_name} Architecture"
228
+ details = f"Dropout: {dropout:.2f}\nBatch Size: {batch_size}"
229
+ ax.set_title(title, fontsize=12, pad=20)
230
+ ax.text(0.5, 0.05, details, ha='center', va='bottom', fontsize=10, transform=ax.transAxes,
231
+ bbox=dict(facecolor='white', alpha=0.8, edgecolor='black'))
232
+
233
  plt.tight_layout()
234
  return fig