DEV Community

Rikin Patel
Rikin Patel

Posted on

Explainable Causal Reinforcement Learning for wildfire evacuation logistics networks with inverse simulation verification

Wildfire Evacuation and AI

Explainable Causal Reinforcement Learning for wildfire evacuation logistics networks with inverse simulation verification

The Spark: A Learning Journey into Crisis AI

It started with a simulation that refused to behave. I was tinkering with a simple evacuation model—a grid of roads, a spreading fire, and a few thousand virtual agents trying to escape. The reinforcement learning agent I'd trained kept sending evacuees toward the fire. Not because it was broken, but because my reward function had inadvertently rewarded "movement" over "survival."

That failure was my gateway into a deeper question: how do we build AI systems that not only optimize complex logistics under extreme uncertainty but also explain why they make those decisions? And how do we verify those explanations when we can't run real-world experiments?

This article chronicles my exploration of Explainable Causal Reinforcement Learning (XC-RL) applied to the nightmarish complexity of wildfire evacuation logistics—and the inverse simulation framework I built to verify what the models claim to have learned.

Why Wildfire Evacuation is the Perfect AI Stress Test

Wildfire evacuation is a "wicked problem" for AI. It combines:

  • Non-stationary dynamics: Fire fronts move unpredictably, wind shifts constantly
  • Strict temporal constraints: Every second of delay compounds risk exponentially
  • Resource scarcity: Limited vehicles, fuel, road capacity, and personnel
  • Human irrationality: Panic, hesitation, and herd behavior defy classical models
  • Causal complexity: Correlation ≠ causation in life-or-death scenarios

Traditional reinforcement learning approaches—like DQN or PPO—can optimize routes, but they operate as opaque black boxes. When an agent reroutes a bus convoy through a smoke-filled canyon, we need to know why. Is it because the canyon is genuinely safer, or because the agent learned a spurious correlation with "shorter distance"?

During my research of causal inference in RL, I discovered that the standard toolkit (SHAP values, LIME, attention maps) provides attributions, not explanations. They tell you which inputs mattered, but not how the system would behave if those inputs were causally intervened upon. For evacuation logistics, that distinction is existential.

The Architecture: Causal RL with an Explanation Layer

My journey led me to a three-tier architecture that I believe represents the future of safety-critical RL:

  1. Causal World Model: A learned graph that encodes the true causal structure of evacuation dynamics (fire spread, road closure cascades, population movement)
  2. Policy Network: A PPO-based agent that acts through the causal model rather than directly from raw observations
  3. Inverse Simulation Verifier: A separate system that runs "counterfactual rollouts" to test whether the policy's stated reasons for actions match actual behavior

Let me walk you through each component as I built it.

Tier 1: The Causal World Model

The first challenge was representing the environment causally. I moved beyond simple state-action pairs to a Structural Causal Model (SCM) with a graph structure G = (V, E), where nodes represent:

  • Fire intensity at location (x, y) at time t
  • Road capacity (dynamic, affected by debris and smoke)
  • Population density in each evacuation zone
  • Available vehicle resources
  • Wind vector and humidity

Edges represent causal relationships. For instance, Fire(x,y,t) → RoadCapacity(x,y,t+1) captures the causal link between fire proximity and road usability.

import networkx as nx
import numpy as np
from typing import Dict, Tuple

class CausalEvacuationModel:
    def __init__(self, road_network: nx.Graph, fire_sources: list):
        self.road_network = road_network
        self.causal_graph = nx.DiGraph()
        self._build_causal_structure(fire_sources)

    def _build_causal_structure(self, fire_sources):
        # Add fire nodes
        for i, source in enumerate(fire_sources):
            self.causal_graph.add_node(f"fire_{i}",
                                      type="fire",
                                      pos=source,
                                      intensity=0.8)

        # Add road nodes with causal edges from fire
        for (u, v, data) in self.road_network.edges(data=True):
            node_id = f"road_{u}_{v}"
            self.causal_graph.add_node(node_id, type="road", capacity=data['capacity'])

            # Causal edge: fire → road (fire reduces capacity)
            for i, source in enumerate(fire_sources):
                dist = np.linalg.norm(np.array(source) - np.array(data['midpoint']))
                if dist < 5000:  # 5km causal radius
                    self.causal_graph.add_edge(f"fire_{i}", node_id,
                                              effect="capacity_reduction",
                                              strength=1.0 / (1.0 + dist/1000))

    def causal_intervention(self, node: str, value: float) -> Dict[str, float]:
        """Do-calculus style intervention: set node value and propagate."""
        # This is where we use Pearl's do-operator
        intervened = {}
        intervened[node] = value

        # Propagate through causal graph (simplified linear approximation)
        for descendant in nx.descendants(self.causal_graph, node):
            # Apply causal effect along edges
            parents = list(self.causal_graph.predecessors(descendant))
            effect = sum(self.causal_graph[parent][descendant]['strength'] *
                        intervened.get(parent, 0.0)
                        for parent in parents if parent in intervened)
            intervened[descendant] = effect

        return intervened
Enter fullscreen mode Exit fullscreen mode

Learning Insight: The critical realization here was that standard do-calculus requires knowing the full causal graph a priori. In real wildfire scenarios, we don't. I had to implement a causal discovery layer using PC algorithm variants on historical fire data, then refine with expert knowledge from fire engineers.

Tier 2: Policy Optimization Through the Causal Lens

The policy network doesn't see raw pixels or sensor readings. Instead, it receives the causal embeddings—representations of the current state that have been transformed through the causal graph. This forces the agent to reason about causes rather than correlations.

import torch
import torch.nn as nn
import torch.optim as optim
from torch.distributions import Categorical

class CausalPolicyNetwork(nn.Module):
    def __init__(self, causal_dim: int, action_dim: int, hidden_dim: int = 128):
        super().__init__()
        self.causal_encoder = nn.Sequential(
            nn.Linear(causal_dim, hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, hidden_dim),
            nn.ReLU()
        )

        # Separate heads for action and explanation
        self.action_head = nn.Linear(hidden_dim, action_dim)
        self.explanation_head = nn.Linear(hidden_dim, causal_dim)  # Predicts causal factors

    def forward(self, causal_state: torch.Tensor):
        encoded = self.causal_encoder(causal_state)
        action_logits = self.action_head(encoded)

        # The explanation head predicts which causal factors drove this decision
        # This is our "explainability" hook
        explanation_weights = torch.softmax(self.explanation_head(encoded), dim=-1)

        return action_logits, explanation_weights

class CausalPPOAgent:
    def __init__(self, causal_model: CausalEvacuationModel):
        self.causal_model = causal_model
        self.policy = CausalPolicyNetwork(
            causal_dim=causal_model.causal_graph.number_of_nodes(),
            action_dim=4  # [north, south, east, west] routing decisions
        )
        self.optimizer = optim.Adam(self.policy.parameters(), lr=3e-4)

    def select_action(self, state: np.ndarray, explain: bool = False):
        # Transform raw state through causal model
        causal_state = self._to_causal_embedding(state)
        causal_tensor = torch.FloatTensor(causal_state).unsqueeze(0)

        action_logits, explanation_weights = self.policy(causal_tensor)

        if explain:
            # Return both action and causal explanation
            action = Categorical(logits=action_logits).sample()
            explanation = explanation_weights.squeeze().detach().numpy()

            # Map explanation weights back to causal graph nodes
            node_names = list(self.causal_model.causal_graph.nodes())
            explanation_dict = {node_names[i]: float(explanation[i])
                              for i in range(len(node_names))}
            return action.item(), explanation_dict

        return Categorical(logits=action_logits).sample().item()
Enter fullscreen mode Exit fullscreen mode

Key Discovery: During my experimentation with this architecture, I found that forcing the policy to output explanation weights before the action (as a kind of "reasoning scratchpad") dramatically improved both performance and interpretability. The agent couldn't "cheat" by making decisions without committing to causal reasons.

Tier 3: Inverse Simulation Verification

This is the part I'm most proud of. The problem with RL explanations is that they're post-hoc—the agent acts, then we try to explain. But what if we could verify the explanation by running the policy backward?

Inverse simulation works like this:

  1. Take a policy decision: "Route Bus 7 through Canyon Road because fire intensity at Ridge Road exceeds threshold"
  2. Run a counterfactual simulation where we intervene on the causal model: set fire intensity at Ridge Road to below threshold
  3. Check if the policy changes its decision
  4. If yes → the explanation is causally valid
  5. If no → the explanation is spurious
import simpy
from typing import List, Dict, Any

class InverseSimulationVerifier:
    def __init__(self, causal_model: CausalEvacuationModel,
                 policy_agent: CausalPPOAgent):
        self.causal_model = causal_model
        self.agent = policy_agent

    def verify_explanation(self, original_state: Dict,
                          original_action: int,
                          explanation: Dict[str, float],
                          num_counterfactuals: int = 100) -> Dict[str, float]:
        """
        Verify if the explanation is causally valid by running
        inverse simulations with counterfactual interventions.
        """
        verification_results = {}

        # For each causal factor in the explanation
        for factor, weight in explanation.items():
            if weight < 0.1:  # Skip negligible factors
                continue

            # Run counterfactual: intervene on this factor
            intervention_value = original_state[factor] * 0.5  # Halve it

            # Run multiple counterfactual simulations
            action_changes = 0
            total_sims = 0

            for _ in range(num_counterfactuals):
                # Create counterfactual state
                cf_state = self.causal_model.causal_intervention(factor, intervention_value)

                # Get policy action under counterfactual
                cf_action = self.agent.select_action(cf_state)

                if cf_action != original_action:
                    action_changes += 1
                total_sims += 1

            # Causal validity score: how often does changing this factor
            # actually change the policy's decision?
            causal_validity = action_changes / total_sims

            verification_results[factor] = {
                'causal_validity': causal_validity,
                'explained_weight': weight,
                'verified': causal_validity > 0.7  # Threshold
            }

        return verification_results

    def generate_verified_explanation(self, state: Dict,
                                     action: int) -> Dict[str, Any]:
        """Generate and verify explanations in one pass."""
        # Get initial explanation from policy
        _, raw_explanation = self.agent.select_action(state, explain=True)

        # Verify through inverse simulation
        verified = self.verify_explanation(state, action, raw_explanation)

        # Filter to only verified causal factors
        verified_factors = {
            factor: data for factor, data in verified.items()
            if data['verified']
        }

        # Build human-readable explanation
        explanation_text = self._format_explanation(verified_factors)

        return {
            'action': action,
            'verified_causal_factors': verified_factors,
            'explanation_text': explanation_text,
            'confidence': np.mean([d['causal_validity'] for d in verified_factors.values()])
        }
Enter fullscreen mode Exit fullscreen mode

During my investigation of inverse verification, I discovered something counterintuitive: the most confident explanations (high weight in the policy's explanation head) were often the least causally valid. The policy had learned to "explain" decisions using salient but non-causal features—like smoke color instead of fire proximity. The inverse simulation caught these spurious explanations and forced the policy to rely on true causal drivers.

Real-World Implementation: A Case Study

Let me show you a complete implementation for a simplified but realistic scenario: evacuating a small town with two evacuation routes, one mountain road and one coastal road.


python
import numpy as np
import matplotlib.pyplot as plt
from dataclasses import dataclass
from typing import List, Tuple

@dataclass
class EvacuationScenario:
    population: int = 5000
    vehicles: int = 1500
    road_capacity: List[int] = None  # vehicles per minute
    fire_speed: float = 0.5  # km per minute
    wind_direction: float = 45  # degrees

    def __post_init__(self):
        if self.road_capacity is None:
            self.road_capacity = [30, 25]  # mountain, coastal

class WildfireEvacuationSystem:
    def __init__(self, scenario: EvacuationScenario):
        self.scenario = scenario
        self.causal_model = self._build_causal_model()
        self.agent = CausalPPOAgent(self.causal_model)
        self.verifier = InverseSimulationVerifier(self.causal_model, self.agent)

    def _build_causal_model(self) -> CausalEvacuationModel:
        # Build road network graph
        road_network = nx.Graph()
        road_network.add_edge("town", "mountain_pass",
                            capacity=self.scenario.road_capacity[0],
                            midpoint=(10, 20))
        road_network.add_edge("town", "coastal_route",
                            capacity=self.scenario.road_capacity[1],
                            midpoint=(5, 15))
        road_network.add_edge("mountain_pass", "safe_zone",
                            capacity=self.scenario.road_capacity[0],
                            midpoint=(20, 25))
        road_network.add_edge("coastal_route", "safe_zone",
                            capacity=self.scenario.road_capacity[1],
                            midpoint=(10, 10))

        # Fire source near coastal route
        fire_sources = [(8, 12)]

        return CausalEvacuationModel(road_network, fire_sources)

    def run_evacuation(self, timesteps: int = 100) -> Dict[str, Any]:
        """Run the full evacuation with explainable decisions."""
        state = self._initialize_state()
        decisions_log = []

        for t in range(timesteps):
            # Agent makes decision with explanation
            action, explanation = self.agent.select_action(state, explain=True)

            # Verify the explanation through inverse simulation
            verified = self.verifier.verify_explanation(state, action, explanation)

            # Log decision with verification
            decisions_log.append({
                'timestep': t,
                'action': action,
                'explanation': verified,
                'state': state.copy()
            })

            # Update state based on action and fire dynamics
            state = self._update_state(state, action, t)

            # Emergency stop if fire reaches town
            if state['fire_proximity'] < 1.0:
                print(f"EVACUATION COMPLETE at timestep {t}")
                break

        return {
            'decisions': decisions_log,
            'total_evacuated': state['evacuated'],
            'casualties': state['population'] - state['evacuated']
        }

    def _update_state(self, state: Dict, action: int, t: int) -> Dict:
        """Update evacuation state based on action and fire dynamics."""
        # Action 0: use mountain route (safer but slower)
        # Action 1: use coastal route (faster but fire risk)
        # Action 2: hold position
        # Action 3: split evacuation

        new_state = state.copy()
        fire_advance = self.scenario.fire_speed * np.cos(np.radians(self.scenario.wind_direction))

        # Update fire position
        new_state['fire_proximity'] -= fire_advance

        if action == 0:
            evacuated = min(state['road_capacity'][0], state['population'])
            new_state['evacuated'] += evacuated
            new_state['population'] -= evacuated
        elif action == 1:
            # Coastal route is faster but riskier
            risk_factor = 1.0 / (state['fire_proximity'] + 0.1)
            evacuated = min(state['road_capacity'][1] * risk_factor, state['population'])
            new_state['evacuated'] += evacuated
            new_state['population'] -= evacuated
            # Fire risk causes casualties
            casualties = int(evacuated * 0.1 * risk_factor)
            new_state['evacuated'] -= casualties
            new_state['casualties'] += casualties
        elif action == 3:
            # Split evacuation
            mountain_cap = state['road_capacity'][0] * 0.6
            coastal_cap = state['road_capacity'][
Enter fullscreen mode Exit fullscreen mode

Top comments (0)