Crash-Intelligence / utils /visualizations.py
shrut27's picture
Upload folder using huggingface_hub
f5823da verified
Raw
History Blame Contribute Delete
9.4 kB
"""Plotly visualization helpers for Crash Intelligence."""
from __future__ import annotations
import pandas as pd
import plotly.express as px
import plotly.graph_objects as go
COLOR_SEQUENCE = [
"#0B6E4F",
"#08A045",
"#1B4965",
"#5FA8D3",
"#C44536",
"#E8871E",
"#6B4C9A",
"#2A9D8F",
"#E76F51",
"#264653",
"#F4A261",
"#457B9D",
"#9B2226",
"#005F73",
"#CA6702",
]
LAYOUT_DEFAULTS = dict(
paper_bgcolor="rgba(255,255,255,1)",
plot_bgcolor="rgba(248,250,252,1)",
font=dict(family="Source Sans 3, Segoe UI, sans-serif", color="#1a1a1a", size=13),
margin=dict(l=50, r=30, t=50, b=50),
legend=dict(bgcolor="rgba(255,255,255,0.9)", bordercolor="#ddd", borderwidth=1),
)
def _apply_layout(fig: go.Figure, title: str, height: int = 420) -> go.Figure:
fig.update_layout(title=title, height=height, **LAYOUT_DEFAULTS)
fig.update_xaxes(showgrid=True, gridcolor="#e5e7eb", zeroline=False)
fig.update_yaxes(showgrid=True, gridcolor="#e5e7eb", zeroline=False)
return fig
def family_bar(summary: pd.DataFrame, metric: str, title: str) -> go.Figure:
fig = px.bar(
summary.sort_values(metric, ascending=True),
x=metric,
y="family",
orientation="h",
color=metric,
color_continuous_scale=["#D8F3DC", "#0B6E4F"],
labels={"family": "Material Family", metric: metric.replace("_", " ").title()},
)
return _apply_layout(fig, title, height=480)
def scatter_crash_vs_weight(materials: pd.DataFrame) -> go.Figure:
fig = px.scatter(
materials.sample(n=min(2000, len(materials)), random_state=3),
x="lightweighting_score",
y="crashworthiness_index",
color="family",
size="uts_mpa",
hover_data=["material_name", "cost_usd_kg", "sustainability_score"],
color_discrete_sequence=COLOR_SEQUENCE,
labels={
"lightweighting_score": "Lightweighting Score",
"crashworthiness_index": "Crashworthiness Index",
},
)
return _apply_layout(fig, "Crashworthiness vs Lightweighting", height=480)
def radar_materials(top: pd.DataFrame) -> go.Figure:
categories = [
"crashworthiness_index",
"lightweighting_score",
"cost_performance_score",
"sustainability_score",
"energy_absorption_potential",
]
labels = ["Crash", "Weight", "Cost-Perf", "Sustainability", "Energy Abs."]
fig = go.Figure()
for i, (_, row) in enumerate(top.head(5).iterrows()):
values = []
for c in categories:
v = float(row[c])
if c == "energy_absorption_potential":
v = min(v * 2.0, 100)
if c == "cost_performance_score":
v = min(v * 1.5, 100)
values.append(v)
values.append(values[0])
fig.add_trace(
go.Scatterpolar(
r=values,
theta=labels + [labels[0]],
name=str(row.get("material_name", row.get("family", f"M{i}"))),
line=dict(color=COLOR_SEQUENCE[i % len(COLOR_SEQUENCE)], width=2),
fill="toself",
opacity=0.55,
)
)
fig.update_layout(
polar=dict(
bgcolor="#f8fafc",
radialaxis=dict(visible=True, range=[0, 100], gridcolor="#e5e7eb"),
angularaxis=dict(gridcolor="#e5e7eb"),
),
title="Multi-Objective Material Comparison",
height=480,
**{k: v for k, v in LAYOUT_DEFAULTS.items() if k != "margin"},
margin=dict(l=60, r=60, t=50, b=40),
)
return fig
def stress_strain_curves(curves: pd.DataFrame, material_ids: list[str]) -> go.Figure:
fig = go.Figure()
subset = curves[curves["material_id"].isin(material_ids)]
for i, mid in enumerate(material_ids):
mdf = subset[subset["material_id"] == mid]
if mdf.empty:
continue
name = mdf["material_name"].iloc[0]
fig.add_trace(
go.Scatter(
x=mdf["strain"],
y=mdf["stress_mpa"],
mode="lines",
name=f"{name} (quasi-static)",
line=dict(color=COLOR_SEQUENCE[i % len(COLOR_SEQUENCE)], width=2.5),
)
)
fig.add_trace(
go.Scatter(
x=mdf["strain"],
y=mdf["stress_high_rate_mpa"],
mode="lines",
name=f"{name} (high-rate)",
line=dict(
color=COLOR_SEQUENCE[i % len(COLOR_SEQUENCE)],
width=2,
dash="dash",
),
)
)
fig.update_layout(
xaxis_title="True Strain",
yaxis_title="True Stress (MPa)",
)
return _apply_layout(fig, "Stress–Strain Curves (Strain-Rate Sensitive)", height=460)
def scenario_heatmap(recommendations: pd.DataFrame) -> go.Figure:
pivot = (
recommendations.groupby(["crash_scenario", "family"])["crash_score"]
.mean()
.reset_index()
.pivot(index="crash_scenario", columns="family", values="crash_score")
)
fig = px.imshow(
pivot,
color_continuous_scale=["#F1FAEE", "#1B4965", "#0B6E4F"],
aspect="auto",
labels=dict(color="Crash Score"),
)
return _apply_layout(fig, "Avg Crash Score by Scenario × Family", height=520)
def energy_intrusion_scatter(recommendations: pd.DataFrame) -> go.Figure:
sample = recommendations.sample(n=min(1500, len(recommendations)), random_state=5)
fig = px.scatter(
sample,
x="intrusion_mm",
y="energy_absorption_kj",
color="crash_scenario",
symbol="family",
hover_data=["material_name", "component", "crash_score"],
color_discrete_sequence=COLOR_SEQUENCE,
labels={
"intrusion_mm": "Intrusion (mm)",
"energy_absorption_kj": "Energy Absorption (kJ)",
},
)
return _apply_layout(fig, "Energy Absorption vs Intrusion", height=460)
def validation_parity(validation: pd.DataFrame) -> go.Figure:
fig = go.Figure()
fig.add_trace(
go.Scatter(
x=validation["cae_crash_score"],
y=validation["ai_crash_score"],
mode="markers",
name="AI vs CAE",
marker=dict(color="#1B4965", size=7, opacity=0.55),
)
)
lims = [0, 100]
fig.add_trace(
go.Scatter(
x=lims,
y=lims,
mode="lines",
name="Ideal",
line=dict(color="#C44536", dash="dash", width=2),
)
)
fig.update_layout(xaxis_title="CAE Crash Score", yaxis_title="AI Crash Score")
return _apply_layout(fig, "AI Prediction vs CAE Validation", height=440)
def validation_error_hist(validation: pd.DataFrame) -> go.Figure:
fig = px.histogram(
validation,
x="ai_cae_error_pct",
nbins=30,
color="pass_fail",
color_discrete_map={"Pass": "#0B6E4F", "Review": "#C44536"},
labels={"ai_cae_error_pct": "AI–CAE Error (%)"},
)
return _apply_layout(fig, "AI–CAE Error Distribution", height=400)
def cost_sustain_bubble(materials: pd.DataFrame) -> go.Figure:
sample = materials.sample(n=min(1500, len(materials)), random_state=9)
fig = px.scatter(
sample,
x="cost_usd_kg",
y="sustainability_score",
size="crashworthiness_index",
color="family",
hover_data=["material_name", "density_g_cm3", "uts_mpa"],
color_discrete_sequence=COLOR_SEQUENCE,
labels={
"cost_usd_kg": "Cost (USD/kg)",
"sustainability_score": "Sustainability Score",
},
)
return _apply_layout(fig, "Cost vs Sustainability (bubble = crash score)", height=460)
def top_recommendations_bar(top: pd.DataFrame) -> go.Figure:
plot_df = top.copy()
name_col = "material_name" if "material_name" in plot_df.columns else "family"
fig = px.bar(
plot_df.sort_values("crash_score", ascending=True),
x="crash_score",
y=name_col,
color="family" if "family" in plot_df.columns else None,
orientation="h",
color_discrete_sequence=COLOR_SEQUENCE,
labels={"crash_score": "Crash Score", name_col: "Material"},
)
return _apply_layout(fig, "Top Recommended Materials", height=420)
def kpi_gauge(value: float, title: str, color: str = "#0B6E4F") -> go.Figure:
fig = go.Figure(
go.Indicator(
mode="gauge+number",
value=value,
title={"text": title, "font": {"size": 14, "color": "#1a1a1a"}},
number={"font": {"color": "#1a1a1a"}},
gauge={
"axis": {"range": [0, 100], "tickcolor": "#1a1a1a"},
"bar": {"color": color},
"bgcolor": "#f1f5f9",
"bordercolor": "#cbd5e1",
"steps": [
{"range": [0, 40], "color": "#fee2e2"},
{"range": [40, 70], "color": "#fef3c7"},
{"range": [70, 100], "color": "#dcfce7"},
],
},
)
)
fig.update_layout(
height=220,
margin=dict(l=20, r=20, t=40, b=10),
paper_bgcolor="white",
font=dict(color="#1a1a1a"),
)
return fig