All checks were successful
ENGIN CI / Build, test and smoke (push) Successful in 43s
188 lines
5.8 KiB
Python
188 lines
5.8 KiB
Python
"""Pure Plotly chart factories, independently testable from Streamlit."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from collections.abc import Sequence
|
|
|
|
import numpy as np
|
|
import plotly.graph_objects as go
|
|
|
|
from .config import FREQ_COLS, LABEL_COLORS, LABEL_DISPLAY
|
|
from .explainability import EngineAnalysis
|
|
|
|
PLOT_BG = "#101820"
|
|
GRID = "rgba(148, 163, 184, 0.14)"
|
|
TEXT = "#dce8ee"
|
|
MUTED = "#8fa6b2"
|
|
COMPARISON_COLORS = ("#28b7d9", "#f5a524", "#38d996", "#d56cff", "#2f78ed")
|
|
HEATMAP_DEVIATION_LIMIT_MV = 20.0
|
|
|
|
|
|
def _base_layout(
|
|
fig: go.Figure,
|
|
*,
|
|
height: int,
|
|
margin_top: int = 42,
|
|
margin_bottom: int = 28,
|
|
) -> go.Figure:
|
|
fig.update_layout(
|
|
template=None,
|
|
height=height,
|
|
margin=dict(l=24, r=24, t=margin_top, b=margin_bottom),
|
|
paper_bgcolor="rgba(0,0,0,0)",
|
|
plot_bgcolor=PLOT_BG,
|
|
font=dict(color=TEXT, family="Inter, system-ui, sans-serif"),
|
|
legend=dict(bgcolor="rgba(0,0,0,0)", font=dict(color=TEXT)),
|
|
hoverlabel=dict(
|
|
bgcolor="#16232c",
|
|
bordercolor="#36505d",
|
|
font_color="#ffffff",
|
|
),
|
|
)
|
|
axis_style = dict(
|
|
color=TEXT,
|
|
gridcolor=GRID,
|
|
linecolor=GRID,
|
|
tickfont=dict(color=TEXT),
|
|
title_font=dict(color=TEXT),
|
|
zeroline=False,
|
|
)
|
|
fig.update_xaxes(**axis_style)
|
|
fig.update_yaxes(**axis_style)
|
|
return fig
|
|
|
|
|
|
def engine_heatmap(analysis: EngineAnalysis) -> go.Figure:
|
|
row_labels = [
|
|
f"C{int(row.cylinder):02d} · {LABEL_DISPLAY[str(row.label)]}"
|
|
for row in analysis.diagnostics.itertuples()
|
|
]
|
|
fig = go.Figure(
|
|
go.Heatmap(
|
|
z=analysis.deviation,
|
|
x=list(range(len(FREQ_COLS))),
|
|
y=row_labels,
|
|
zmin=-HEATMAP_DEVIATION_LIMIT_MV,
|
|
zmax=HEATMAP_DEVIATION_LIMIT_MV,
|
|
zmid=0,
|
|
colorscale=[
|
|
[0.0, "#28b7d9"],
|
|
[0.5, "#13222b"],
|
|
[1.0, "#ff6b57"],
|
|
],
|
|
colorbar=dict(
|
|
title=dict(text="Δ mV", font=dict(color=TEXT)),
|
|
tickfont=dict(color=TEXT),
|
|
bgcolor="rgba(0,0,0,0)",
|
|
borderwidth=0,
|
|
thickness=12,
|
|
),
|
|
hovertemplate="%{y}<br>%{x} kHz<br>Odchylenie %{z:.1f} mV<extra></extra>",
|
|
)
|
|
)
|
|
fig.update_yaxes(autorange="reversed", automargin=True, title=None)
|
|
fig.update_xaxes(title="Częstotliwość [kHz]", dtick=2)
|
|
return _base_layout(fig, height=360, margin_top=18, margin_bottom=42)
|
|
|
|
|
|
def _cylinder_position(analysis: EngineAnalysis, cylinder: int) -> int:
|
|
positions = np.flatnonzero(analysis.measurements["cylinder"].to_numpy() == cylinder)
|
|
if len(positions) != 1:
|
|
raise KeyError(f"Unknown cylinder={cylinder}")
|
|
return int(positions[0])
|
|
|
|
|
|
def cylinder_spectrum(
|
|
analysis: EngineAnalysis,
|
|
cylinder: int,
|
|
comparison_cylinders: Sequence[int] | None = None,
|
|
) -> go.Figure:
|
|
selected = list(
|
|
dict.fromkeys([cylinder, *(comparison_cylinders or ())])
|
|
)
|
|
if len(selected) > 4:
|
|
raise ValueError("At most four cylinders can be compared.")
|
|
|
|
position = _cylinder_position(analysis, cylinder)
|
|
row = analysis.diagnostics.iloc[position]
|
|
primary_color = LABEL_COLORS[str(row["label"])]
|
|
frequency = np.arange(len(FREQ_COLS))
|
|
fig = go.Figure()
|
|
fig.add_trace(
|
|
go.Scatter(
|
|
x=frequency,
|
|
y=analysis.reference[position],
|
|
mode="lines",
|
|
name=f"Referencja C{cylinder:02d}",
|
|
line=dict(color="#7f95a1", width=2, dash="dash"),
|
|
)
|
|
)
|
|
|
|
extra_colors = iter(
|
|
color
|
|
for color in COMPARISON_COLORS
|
|
if color.lower() != primary_color.lower()
|
|
)
|
|
for selected_cylinder in selected:
|
|
selected_position = _cylinder_position(analysis, selected_cylinder)
|
|
color = primary_color if selected_cylinder == cylinder else next(extra_colors)
|
|
fig.add_trace(
|
|
go.Scatter(
|
|
x=frequency,
|
|
y=analysis.spectra[selected_position],
|
|
mode="lines+markers",
|
|
name=f"C{selected_cylinder:02d}",
|
|
line=dict(
|
|
color=color,
|
|
width=3 if selected_cylinder == cylinder else 2,
|
|
),
|
|
marker=dict(size=6 if selected_cylinder == cylinder else 4),
|
|
)
|
|
)
|
|
|
|
top = {
|
|
int(value)
|
|
for value in str(row["top_anomalous_frequencies_khz"]).split("|")
|
|
if value
|
|
}
|
|
for value in top:
|
|
fig.add_vrect(
|
|
x0=value - 0.35,
|
|
x1=value + 0.35,
|
|
fillcolor=primary_color,
|
|
opacity=0.10,
|
|
line_width=0,
|
|
)
|
|
fig.update_layout(
|
|
legend=dict(
|
|
orientation="h",
|
|
x=0,
|
|
xanchor="left",
|
|
y=1.04,
|
|
yanchor="bottom",
|
|
)
|
|
)
|
|
fig.update_xaxes(title="Częstotliwość [kHz]", dtick=1)
|
|
fig.update_yaxes(title="Amplituda [mV]")
|
|
return _base_layout(fig, height=440, margin_top=76, margin_bottom=42)
|
|
|
|
|
|
def deviation_chart(analysis: EngineAnalysis, cylinder: int) -> go.Figure:
|
|
position = int(
|
|
np.flatnonzero(analysis.measurements["cylinder"].to_numpy() == cylinder)[0]
|
|
)
|
|
values = analysis.deviation[position]
|
|
colors = ["#ff6b57" if value >= 0 else "#28b7d9" for value in values]
|
|
fig = go.Figure(
|
|
go.Bar(
|
|
x=np.arange(len(FREQ_COLS)),
|
|
y=values,
|
|
marker_color=colors,
|
|
hovertemplate="%{x} kHz<br>Δ %{y:.1f} mV<extra></extra>",
|
|
)
|
|
)
|
|
fig.add_hline(y=0, line_color="#6e8794", line_width=1)
|
|
fig.update_xaxes(title="Częstotliwość [kHz]", dtick=2)
|
|
fig.update_yaxes(title="Różnica [mV]")
|
|
return _base_layout(fig, height=340, margin_top=24, margin_bottom=42)
|