Spaces:
Runtime error
Runtime error
Update core/plot.py
Browse files- 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 ['
|
| 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=
|
| 73 |
|
| 74 |
-
|
| 75 |
-
max_val = max(metrics.values(), default=1)
|
| 76 |
min_val = min(metrics.values(), default=0)
|
| 77 |
-
ax.set_ylim(min(min_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:.
|
| 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
|