Skip to content

Core: Data

All dataset, preprocessing, representation, and discovery machinery.

Current layout

  • datasets/ - raw source adapters and dataset/source builders.
  • discovery/ - signal profiles, canonical entities, and provisional hypotheses for cross-vehicle alignment.
  • preprocessing/ - explicit representations, views, segments, materialization, PyG packing, temporal streams, scaler config, vocab config, and graph transforms.
  • datamodule/ - training-time loaders and batching policy.
  • state.py - process-local dataset state for reuse within one Python process.

graphids.core.data

data

Data-layer public API.

The runtime datamodules are imported lazily so preprocessing and discovery can be used without importing optional training dependencies.

CANBusTemporalSource dataclass

CANBusTemporalSource(name: str, lake_root: str | None = None, val_fraction: float = 0.2, representation_cfg: TemporalRepresentationCfg = TemporalRepresentationCfg(), vocab_scope: str = 'train', train_source_mode: str = 'mixed', val_warmup_events: int = 0, test_warmup_events: int = 0)

Catalog to train/val/test CAN TemporalData cache builder.

DatasetState dataclass

DatasetState(train: Any, val: Any, test: dict[str, Any])

Ready-to-serve train/val/test splits.

clear_cache

clear_cache() -> None

Drop all cached states. Intended for test teardown.

Source code in graphids/core/data/state.py
def clear_cache() -> None:
    """Drop all cached states. Intended for test teardown."""
    _REGISTRY.clear()

get_or_build

get_or_build(dataset: _CacheableDataset) -> DatasetState

Return cached DatasetState for dataset.

Source code in graphids/core/data/state.py
def get_or_build(dataset: _CacheableDataset) -> DatasetState:
    """Return cached ``DatasetState`` for ``dataset``."""
    key = dataset.cache_key
    state = _REGISTRY.get(key)
    if state is None:
        state = dataset.build()
        _REGISTRY[key] = state
    return state

datamodule

DataModule primitives for temporal datasets.

TemporalDataModule

TemporalDataModule(dataset, batch_size: int = 256, *, batch_mode: str = 'events', stream_lanes: int | None = None, chunk_size: int | None = None, num_workers: int = 0, pin_memory: bool = False, persistent_workers: bool = False)

Bases: LightningDataModule

Serve temporal event streams with PyG's TemporalDataLoader.

Source code in graphids/core/data/datamodule/temporal.py
def __init__(
    self,
    dataset,
    batch_size: int = 256,
    *,
    batch_mode: str = "events",
    stream_lanes: int | None = None,
    chunk_size: int | None = None,
    num_workers: int = 0,
    pin_memory: bool = False,
    persistent_workers: bool = False,
):
    super().__init__()
    self.source = dataset
    self.batch_size = batch_size
    self.batch_mode = str(batch_mode)
    self.stream_lanes = None if stream_lanes is None else int(stream_lanes)
    self.chunk_size = None if chunk_size is None else int(chunk_size)
    self.num_workers = int(num_workers)
    self.pin_memory = bool(pin_memory)
    self.persistent_workers = bool(persistent_workers)
    if self.batch_mode not in {"events", "stream_lanes"}:
        raise ValueError("batch_mode must be one of: events, stream_lanes")
    if self.batch_mode == "stream_lanes":
        if self.stream_lanes is None or self.stream_lanes < 2:
            raise ValueError("stream_lanes must be >= 2 for batch_mode='stream_lanes'")
        if self.chunk_size is None or self.chunk_size < 1:
            raise ValueError("chunk_size must be positive for batch_mode='stream_lanes'")
    self._train = None
    self._val = None
    self._tests: dict[str, object] = {}
num_ids property
num_ids: int

Embedding table size for temporal source/destination ids.

test_datasets property
test_datasets: dict[str, object]

Compatibility name used by model test-set bookkeeping.

stream_lanes

Lane-batched temporal event iteration.

TemporalStreamLaneBatch dataclass
TemporalStreamLaneBatch(src: Tensor, dst: Tensor, t: Tensor, msg: Tensor, y: Tensor, attack_type: Tensor, event_id: Tensor, is_scored: Tensor, stream_id: Tensor, valid_mask: Tensor, lane_reset: Tensor, lane_stream_end: Tensor, reset_after: Tensor)

A fixed-lane batch of independent temporal streams.

Event tensors are shaped [lanes, chunk_size, ...]. valid_mask marks real events; padded slots carry neutral values and must not affect loss or metrics.

TemporalStreamLaneLoader
TemporalStreamLaneLoader(data, *, stream_lanes: int, chunk_size: int)

Deterministically pack independent streams into fixed temporal lanes.

Source code in graphids/core/data/datamodule/stream_lanes.py
def __init__(self, data, *, stream_lanes: int, chunk_size: int):
    self.data = data
    self.stream_lanes = int(stream_lanes)
    self.chunk_size = int(chunk_size)
    if self.stream_lanes < 2:
        raise ValueError("stream_lanes must be >= 2 for batch_mode='stream_lanes'")
    if self.chunk_size < 1:
        raise ValueError("chunk_size must be positive for batch_mode='stream_lanes'")
    self._streams = self._partition_streams(data)
    self._length = self._compute_length()

temporal

Lightning data module for temporal PyG event streams.

TemporalDataModule
TemporalDataModule(dataset, batch_size: int = 256, *, batch_mode: str = 'events', stream_lanes: int | None = None, chunk_size: int | None = None, num_workers: int = 0, pin_memory: bool = False, persistent_workers: bool = False)

Bases: LightningDataModule

Serve temporal event streams with PyG's TemporalDataLoader.

Source code in graphids/core/data/datamodule/temporal.py
def __init__(
    self,
    dataset,
    batch_size: int = 256,
    *,
    batch_mode: str = "events",
    stream_lanes: int | None = None,
    chunk_size: int | None = None,
    num_workers: int = 0,
    pin_memory: bool = False,
    persistent_workers: bool = False,
):
    super().__init__()
    self.source = dataset
    self.batch_size = batch_size
    self.batch_mode = str(batch_mode)
    self.stream_lanes = None if stream_lanes is None else int(stream_lanes)
    self.chunk_size = None if chunk_size is None else int(chunk_size)
    self.num_workers = int(num_workers)
    self.pin_memory = bool(pin_memory)
    self.persistent_workers = bool(persistent_workers)
    if self.batch_mode not in {"events", "stream_lanes"}:
        raise ValueError("batch_mode must be one of: events, stream_lanes")
    if self.batch_mode == "stream_lanes":
        if self.stream_lanes is None or self.stream_lanes < 2:
            raise ValueError("stream_lanes must be >= 2 for batch_mode='stream_lanes'")
        if self.chunk_size is None or self.chunk_size < 1:
            raise ValueError("chunk_size must be positive for batch_mode='stream_lanes'")
    self._train = None
    self._val = None
    self._tests: dict[str, object] = {}
num_ids property
num_ids: int

Embedding table size for temporal source/destination ids.

test_datasets property
test_datasets: dict[str, object]

Compatibility name used by model test-set bookkeeping.

datasets

CANBusTemporalSource dataclass

CANBusTemporalSource(name: str, lake_root: str | None = None, val_fraction: float = 0.2, representation_cfg: TemporalRepresentationCfg = TemporalRepresentationCfg(), vocab_scope: str = 'train', train_source_mode: str = 'mixed', val_warmup_events: int = 0, test_warmup_events: int = 0)

Catalog to train/val/test CAN TemporalData cache builder.

can_bus

CAN bus temporal dataset adapter and row schema.

CANBusTemporalSource dataclass
CANBusTemporalSource(name: str, lake_root: str | None = None, val_fraction: float = 0.2, representation_cfg: TemporalRepresentationCfg = TemporalRepresentationCfg(), vocab_scope: str = 'train', train_source_mode: str = 'mixed', val_warmup_events: int = 0, test_warmup_events: int = 0)

Catalog to train/val/test CAN TemporalData cache builder.

infer_attack_type
infer_attack_type(csv: Path) -> int

Infer the attack code from filename/path substrings.

Source code in graphids/core/data/datasets/can_bus.py
def infer_attack_type(csv: Path) -> int:
    """Infer the attack code from filename/path substrings."""
    s = csv.stem.lower() + " " + csv.parent.name.lower()
    for kw, code in ATTACK_TYPE_CODES.items():
        if kw in s:
            return code
    return 0
load_can_rows
load_can_rows(raw_dir: Path, source_dirs: list[str]) -> pl.DataFrame

Load, normalize, and parse raw CAN CSVs from source dirs.

Source code in graphids/core/data/datasets/can_bus.py
def load_can_rows(raw_dir: Path, source_dirs: list[str]) -> pl.DataFrame:
    """Load, normalize, and parse raw CAN CSVs from source dirs."""
    if not source_dirs:
        raise ValueError("source_dirs is empty; cannot load CAN rows")
    frames: list[pl.LazyFrame] = []
    for sub in source_dirs:
        sub_path = raw_dir / sub
        if not sub_path.is_dir():
            raise FileNotFoundError(f"declared source_dir {sub!r} missing under {raw_dir}")
        for csv_path in sorted(sub_path.rglob("*.csv")):
            at = infer_attack_type(csv_path)
            frames.append(
                pl.scan_csv(csv_path).with_columns(
                    pl.lit(at).alias("attack_type"),
                    pl.lit(sub).alias("vehicle_id"),
                    pl.lit(sub).alias("source_dir"),
                    pl.lit(str(csv_path.relative_to(raw_dir))).alias("source_file"),
                )
            )
    if not frames:
        raise ValueError(f"no CSVs under any of {source_dirs!r} in {raw_dir}")

    combined = pl.concat(frames).sort("timestamp")
    cols = combined.collect_schema().names()
    renames = {
        old: new
        for old, new in (("arbitration_id", "arb_id"), ("data_field", "payload"))
        if old in cols
    }
    if renames:
        combined = combined.rename(renames)
    return parse_payload(combined).collect()
parse_payload
parse_payload(lf: LazyFrame) -> pl.LazyFrame

Hex payload to byte_0..7 plus Shannon entropy.

Source code in graphids/core/data/datasets/can_bus.py
def parse_payload(lf: pl.LazyFrame) -> pl.LazyFrame:
    """Hex ``payload`` to ``byte_0..7`` plus Shannon entropy."""
    if "byte_0" in lf.collect_schema().names():
        return lf
    byte_exprs = [
        pl.col("payload").cast(pl.Utf8).str.slice(i * 2, 2).str.to_integer(base=16, strict=False)
        .fill_null(0).cast(pl.Float32).alias(f"byte_{i}")
        for i in range(N_BYTES)
    ]
    lf = lf.with_columns(byte_exprs)
    bcols = [pl.col(c) for c in BYTE_COLS]
    row_sum = pl.sum_horizontal(bcols).clip(1e-12, None)
    entropy = pl.sum_horizontal(
        [pl.when(c > 0).then(-(c / row_sum) * (c / row_sum).log()).otherwise(0.0) for c in bcols]
    ).alias("entropy")
    return lf.with_columns(entropy)

discovery

Signal profile artifacts for CAN cache builds.

build_signal_profiles

build_signal_profiles(df: DataFrame) -> pl.DataFrame

Aggregate raw CAN rows into one profile per vehicle/arbitration ID.

Source code in graphids/core/data/discovery/hypotheses.py
def build_signal_profiles(df: pl.DataFrame) -> pl.DataFrame:
    """Aggregate raw CAN rows into one profile per vehicle/arbitration ID."""

    missing = [c for c in ("vehicle_id", "arb_id") if c not in df.columns]
    if missing:
        raise ValueError(f"build_signal_profiles missing columns: {missing}")

    sort_cols = [c for c in ("vehicle_id", "arb_id", "timestamp") if c in df.columns]
    if sort_cols:
        df = df.sort(sort_cols)

    aggs: list[pl.Expr] = [pl.len().cast(pl.Int64).alias("msg_count")]
    if "timestamp" in df.columns:
        aggs.extend(
            [
                pl.col("timestamp").min().cast(pl.Float64).alias("timestamp_min"),
                pl.col("timestamp").max().cast(pl.Float64).alias("timestamp_max"),
                (pl.col("timestamp").max() - pl.col("timestamp").min()).cast(pl.Float64).alias("duration"),
                pl.col("timestamp").diff().mean().cast(pl.Float64).alias("iat_mean"),
                pl.col("timestamp").diff().std().fill_nan(0).cast(pl.Float64).alias("iat_std"),
            ]
        )
    if "entropy" in df.columns:
        aggs.extend(
            [
                pl.col("entropy").mean().cast(pl.Float64).alias("entropy_mean"),
                pl.col("entropy").std().fill_nan(0).cast(pl.Float64).alias("entropy_std"),
            ]
        )

    byte_cols = _byte_cols(df)
    aggs.extend(pl.col(c).mean().cast(pl.Float64).alias(f"{c}_mean") for c in byte_cols)
    aggs.extend(pl.col(c).std().fill_nan(0).cast(pl.Float64).alias(f"{c}_std") for c in byte_cols)
    aggs.extend((pl.col(c).max() - pl.col(c).min()).cast(pl.Float64).alias(f"{c}_range") for c in byte_cols)
    if byte_cols:
        aggs.append(
            pl.mean_horizontal(
                *[(pl.col(c).diff().abs().drop_nulls() > 0).mean().fill_null(0) for c in byte_cols]
            ).cast(pl.Float64).alias("change_rate")
        )
    if "attack" in df.columns:
        aggs.extend(
            [
                pl.col("attack").max().cast(pl.Int64).alias("attack_max"),
                pl.col("attack").mean().cast(pl.Float64).alias("attack_rate"),
            ]
        )

    return df.group_by("vehicle_id", "arb_id").agg(*aggs).with_columns(
        pl.concat_str([pl.col("vehicle_id").cast(pl.Utf8), pl.col("arb_id").cast(pl.Utf8)], separator="::").alias("signal_key")
    )

initialize_hypotheses

initialize_hypotheses(profiles: DataFrame) -> pl.DataFrame

Create empty editable mapping rows for profile review.

Source code in graphids/core/data/discovery/hypotheses.py
def initialize_hypotheses(profiles: pl.DataFrame) -> pl.DataFrame:
    """Create empty editable mapping rows for profile review."""

    required = ["vehicle_id", "arb_id", "signal_key"]
    missing = [c for c in required if c not in profiles.columns]
    if missing:
        raise ValueError(f"initialize_hypotheses missing columns: {missing}")
    return profiles.select(*required).with_columns(
        pl.lit(None, dtype=pl.Utf8).alias("candidate_canonical_id"),
        pl.lit(0.0).cast(pl.Float64).alias("confidence"),
        pl.lit("unreviewed").alias("status"),
        pl.lit("").cast(pl.Utf8).alias("evidence"),
    )

hypotheses

Signal profile artifacts written beside graph caches.

build_signal_profiles
build_signal_profiles(df: DataFrame) -> pl.DataFrame

Aggregate raw CAN rows into one profile per vehicle/arbitration ID.

Source code in graphids/core/data/discovery/hypotheses.py
def build_signal_profiles(df: pl.DataFrame) -> pl.DataFrame:
    """Aggregate raw CAN rows into one profile per vehicle/arbitration ID."""

    missing = [c for c in ("vehicle_id", "arb_id") if c not in df.columns]
    if missing:
        raise ValueError(f"build_signal_profiles missing columns: {missing}")

    sort_cols = [c for c in ("vehicle_id", "arb_id", "timestamp") if c in df.columns]
    if sort_cols:
        df = df.sort(sort_cols)

    aggs: list[pl.Expr] = [pl.len().cast(pl.Int64).alias("msg_count")]
    if "timestamp" in df.columns:
        aggs.extend(
            [
                pl.col("timestamp").min().cast(pl.Float64).alias("timestamp_min"),
                pl.col("timestamp").max().cast(pl.Float64).alias("timestamp_max"),
                (pl.col("timestamp").max() - pl.col("timestamp").min()).cast(pl.Float64).alias("duration"),
                pl.col("timestamp").diff().mean().cast(pl.Float64).alias("iat_mean"),
                pl.col("timestamp").diff().std().fill_nan(0).cast(pl.Float64).alias("iat_std"),
            ]
        )
    if "entropy" in df.columns:
        aggs.extend(
            [
                pl.col("entropy").mean().cast(pl.Float64).alias("entropy_mean"),
                pl.col("entropy").std().fill_nan(0).cast(pl.Float64).alias("entropy_std"),
            ]
        )

    byte_cols = _byte_cols(df)
    aggs.extend(pl.col(c).mean().cast(pl.Float64).alias(f"{c}_mean") for c in byte_cols)
    aggs.extend(pl.col(c).std().fill_nan(0).cast(pl.Float64).alias(f"{c}_std") for c in byte_cols)
    aggs.extend((pl.col(c).max() - pl.col(c).min()).cast(pl.Float64).alias(f"{c}_range") for c in byte_cols)
    if byte_cols:
        aggs.append(
            pl.mean_horizontal(
                *[(pl.col(c).diff().abs().drop_nulls() > 0).mean().fill_null(0) for c in byte_cols]
            ).cast(pl.Float64).alias("change_rate")
        )
    if "attack" in df.columns:
        aggs.extend(
            [
                pl.col("attack").max().cast(pl.Int64).alias("attack_max"),
                pl.col("attack").mean().cast(pl.Float64).alias("attack_rate"),
            ]
        )

    return df.group_by("vehicle_id", "arb_id").agg(*aggs).with_columns(
        pl.concat_str([pl.col("vehicle_id").cast(pl.Utf8), pl.col("arb_id").cast(pl.Utf8)], separator="::").alias("signal_key")
    )
initialize_hypotheses
initialize_hypotheses(profiles: DataFrame) -> pl.DataFrame

Create empty editable mapping rows for profile review.

Source code in graphids/core/data/discovery/hypotheses.py
def initialize_hypotheses(profiles: pl.DataFrame) -> pl.DataFrame:
    """Create empty editable mapping rows for profile review."""

    required = ["vehicle_id", "arb_id", "signal_key"]
    missing = [c for c in required if c not in profiles.columns]
    if missing:
        raise ValueError(f"initialize_hypotheses missing columns: {missing}")
    return profiles.select(*required).with_columns(
        pl.lit(None, dtype=pl.Utf8).alias("candidate_canonical_id"),
        pl.lit(0.0).cast(pl.Float64).alias("confidence"),
        pl.lit("unreviewed").alias("status"),
        pl.lit("").cast(pl.Utf8).alias("evidence"),
    )

preprocessing

Core temporal preprocessing.

add_temporal_split_masks

add_temporal_split_masks(table: DataFrame, *, split_name: str, warmup_events: int = 0) -> pl.DataFrame

Attach split identity plus warmup/scoring masks to an event table.

Source code in graphids/core/data/preprocessing/temporal.py
def add_temporal_split_masks(
    table: pl.DataFrame,
    *,
    split_name: str,
    warmup_events: int = 0,
) -> pl.DataFrame:
    """Attach split identity plus warmup/scoring masks to an event table."""
    if split_name not in SPLIT_NAME_TO_ID:
        raise ValueError(f"unsupported split_name {split_name!r}; expected one of {sorted(SPLIT_NAME_TO_ID)}")
    if warmup_events < 0:
        raise ValueError("warmup_events must be non-negative")
    if table.is_empty():
        return table.with_columns(
            pl.lit(split_name).alias("split_name"),
            pl.lit(SPLIT_NAME_TO_ID[split_name]).cast(pl.Int64).alias("split_id"),
            pl.lit(False).alias("is_warmup"),
            pl.lit(False).alias("is_scored"),
        )
    rank = _rank_within_stream(table)
    return table.with_columns(
        pl.lit(split_name).alias("split_name"),
        pl.lit(SPLIT_NAME_TO_ID[split_name]).cast(pl.Int64).alias("split_id"),
        (rank < warmup_events).alias("is_warmup"),
    ).with_columns((~pl.col("is_warmup")).alias("is_scored"))

assert_temporal_splits_disjoint

assert_temporal_splits_disjoint(*tables: DataFrame) -> None

Raise if any split shares event ids with another split.

Source code in graphids/core/data/preprocessing/temporal.py
def assert_temporal_splits_disjoint(*tables: pl.DataFrame) -> None:
    """Raise if any split shares event ids with another split."""
    seen: set[int] = set()
    for table in tables:
        ids = _event_ids(table)
        overlap = seen & ids
        if overlap:
            sample = sorted(overlap)[:5]
            raise ValueError(f"temporal splits share event_id values: {sample}")
        seen |= ids

build_temporal_event_table

build_temporal_event_table(df: DataFrame, *, id_col: str = 'node_id', raw_id_col: str = 'arb_id', timestamp_col: str = 'timestamp', num_unknown_buckets: int = 128) -> pl.DataFrame

Build one temporal event per row, preserving stream provenance.

The first row in each stream is represented as a self-event (src_id == dst_id) so event counts match normalized row counts.

Source code in graphids/core/data/preprocessing/temporal.py
def build_temporal_event_table(
    df: pl.DataFrame,
    *,
    id_col: str = "node_id",
    raw_id_col: str = "arb_id",
    timestamp_col: str = "timestamp",
    num_unknown_buckets: int = 128,
) -> pl.DataFrame:
    """Build one temporal event per row, preserving stream provenance.

    The first row in each stream is represented as a self-event
    (``src_id == dst_id``) so event counts match normalized row counts.
    """
    if id_col not in df.columns:
        raise ValueError(f"missing mapped id column {id_col!r}")
    if timestamp_col not in df.columns:
        raise ValueError(f"missing timestamp column {timestamp_col!r}")

    df = _ensure_columns(
        df,
        {
            raw_id_col: "",
            "attack": 0,
            "attack_type": 0,
            "entropy": 0.0,
            "vehicle_id": "",
            "source_dir": "",
            "source_file": "",
        },
    )
    df = _ensure_columns(df, {col: 0.0 for col in TEMPORAL_BYTE_COLS})
    stream_keys = ["vehicle_id", "source_dir", "source_file"]
    rows = df.with_row_index("row_index").with_columns(
        pl.col("row_index").cast(pl.Int64),
        pl.col(id_col).cast(pl.Int64).alias("dst_id"),
        pl.col(raw_id_col).cast(pl.Utf8).alias("dst_raw"),
        pl.col(timestamp_col).cast(pl.Float64).alias("timestamp"),
        pl.col("attack").fill_null(0).cast(pl.Int64),
        pl.col("attack_type").fill_null(0).cast(pl.Int64),
        *(pl.col(c).fill_null(0).cast(pl.Float32) for c in TEMPORAL_BYTE_COLS),
        pl.col("entropy").fill_null(0).cast(pl.Float32),
    )
    streams = rows.select(stream_keys).unique(maintain_order=True).with_row_index("stream_id")
    rows = (
        rows.join(streams, on=stream_keys, how="left")
        .with_columns(pl.col("stream_id").cast(pl.Int64))
        .sort(["stream_id", "timestamp", "row_index"])
    )

    delta_exprs = [
        (pl.col(c) - pl.col(c).shift(1).over("stream_id"))
        .fill_null(0)
        .cast(pl.Float32)
        .alias(f"{c}_delta")
        for c in TEMPORAL_BYTE_COLS
    ]
    table = rows.with_columns(
        pl.col("dst_id").shift(1).over("stream_id").fill_null(pl.col("dst_id")).cast(pl.Int64).alias("src_id"),
        pl.col("dst_raw").shift(1).over("stream_id").fill_null(pl.col("dst_raw")).alias("src_raw"),
        (pl.col("timestamp") - pl.col("timestamp").shift(1).over("stream_id"))
        .fill_null(0)
        .cast(pl.Float32)
        .alias("iat"),
        *delta_exprs,
    )
    table = table.with_columns(
        (pl.col("src_id") == 0).alias("src_is_unknown"),
        (pl.col("dst_id") == 0).alias("dst_is_unknown"),
        (
            pl.col("stream_id").shift(-1).is_null()
            | (pl.col("stream_id").shift(-1) != pl.col("stream_id"))
        ).alias("reset_after"),
        (pl.col("attack") > 0).cast(pl.Int64).alias("y"),
    )
    table = _with_unknown_buckets(table, num_unknown_buckets=num_unknown_buckets)
    msg_feature_cols = [
        c
        for c in TEMPORAL_MSG_COL_ORDER
        if c
        not in {
            "src_is_unknown",
            "dst_is_unknown",
            "src_unknown_bucket",
            "dst_unknown_bucket",
        }
    ]
    return table.with_row_index("event_id").with_columns(pl.col("event_id").cast(pl.Int64)).select(
        "event_id",
        "vehicle_id",
        "source_dir",
        "source_file",
        "row_index",
        "timestamp",
        "src_id",
        "dst_id",
        "src_raw",
        "dst_raw",
        "src_is_unknown",
        "dst_is_unknown",
        "src_unknown_bucket",
        "dst_unknown_bucket",
        "stream_id",
        "reset_after",
        *msg_feature_cols,
        "y",
        "attack_type",
    )

prepare_temporal_eval_table

prepare_temporal_eval_table(table: DataFrame, *, split_name: str = 'test', warmup_events: int = 0) -> pl.DataFrame

Prepare validation/test-style full streams with warmup/scoring masks.

Source code in graphids/core/data/preprocessing/temporal.py
def prepare_temporal_eval_table(
    table: pl.DataFrame,
    *,
    split_name: str = "test",
    warmup_events: int = 0,
) -> pl.DataFrame:
    """Prepare validation/test-style full streams with warmup/scoring masks."""
    return add_temporal_split_masks(
        mark_terminal_reset(mark_split_start_self_events(table)),
        split_name=split_name,
        warmup_events=warmup_events,
    )

split_temporal_train_val_tables

split_temporal_train_val_tables(table: DataFrame, *, val_fraction: float, val_warmup_events: int = 0) -> tuple[pl.DataFrame, pl.DataFrame]

Chronologically split each stream into train and validation intervals.

Source code in graphids/core/data/preprocessing/temporal.py
def split_temporal_train_val_tables(
    table: pl.DataFrame,
    *,
    val_fraction: float,
    val_warmup_events: int = 0,
) -> tuple[pl.DataFrame, pl.DataFrame]:
    """Chronologically split each stream into train and validation intervals."""
    train_ids, val_ids = _split_event_ids_by_stream(table, val_fraction=val_fraction)
    train = table.filter(pl.col("event_id").is_in(train_ids))
    val = table.filter(pl.col("event_id").is_in(val_ids))
    train = add_temporal_split_masks(
        mark_terminal_reset(mark_split_start_self_events(train)),
        split_name="train",
        warmup_events=0,
    )
    val = add_temporal_split_masks(
        mark_terminal_reset(mark_split_start_self_events(val)),
        split_name="val",
        warmup_events=val_warmup_events,
    )
    assert_temporal_splits_disjoint(train, val)
    return train, val

temporal_to_pyg

temporal_to_pyg(table: DataFrame) -> TemporalData

Pack a temporal event table into PyG TemporalData tensors.

Source code in graphids/core/data/preprocessing/temporal.py
def temporal_to_pyg(table: pl.DataFrame) -> TemporalData:
    """Pack a temporal event table into PyG ``TemporalData`` tensors."""
    optional_tensors = {}
    optional_specs = {
        "is_warmup": pl.Boolean,
        "is_scored": pl.Boolean,
        "split_id": pl.Int64,
    }
    for col, dtype in optional_specs.items():
        if col in table.columns:
            optional_tensors[col] = _tensor(table, col, dtype=dtype)
    data = TemporalData(
        src=_tensor(table, "src_id", dtype=pl.Int64),
        dst=_tensor(table, "dst_id", dtype=pl.Int64),
        t=_tensor(table, "timestamp", dtype=pl.Float32),
        msg=_tensor(table, list(TEMPORAL_MSG_COL_ORDER), dtype=pl.Float32),
        y=_tensor(table, "y", dtype=pl.Int64),
        attack_type=_tensor(table, "attack_type", dtype=pl.Int64),
        stream_id=_tensor(table, "stream_id", dtype=pl.Int64),
        reset_after=_tensor(table, "reset_after", dtype=pl.Boolean),
        event_id=_tensor(table, "event_id", dtype=pl.Int64),
        **optional_tensors,
    )
    if "split_name" in table.columns:
        names = table["split_name"].unique().to_list()
        # ``TemporalData`` treats normal attributes as event fields and slices
        # them in ``TemporalDataLoader``. Keep this human-readable label as
        # object metadata; the tensor split contract is ``split_id``.
        object.__setattr__(data, "split_name", str(names[0]) if len(names) == 1 else "mixed")
    return data

representations

Data representation configs used by preprocessing.

scaler

Per-column feature scalers for tensor-based graph preprocessing.

temporal

Temporal event materialization for normalized CAN rows.

add_temporal_split_masks
add_temporal_split_masks(table: DataFrame, *, split_name: str, warmup_events: int = 0) -> pl.DataFrame

Attach split identity plus warmup/scoring masks to an event table.

Source code in graphids/core/data/preprocessing/temporal.py
def add_temporal_split_masks(
    table: pl.DataFrame,
    *,
    split_name: str,
    warmup_events: int = 0,
) -> pl.DataFrame:
    """Attach split identity plus warmup/scoring masks to an event table."""
    if split_name not in SPLIT_NAME_TO_ID:
        raise ValueError(f"unsupported split_name {split_name!r}; expected one of {sorted(SPLIT_NAME_TO_ID)}")
    if warmup_events < 0:
        raise ValueError("warmup_events must be non-negative")
    if table.is_empty():
        return table.with_columns(
            pl.lit(split_name).alias("split_name"),
            pl.lit(SPLIT_NAME_TO_ID[split_name]).cast(pl.Int64).alias("split_id"),
            pl.lit(False).alias("is_warmup"),
            pl.lit(False).alias("is_scored"),
        )
    rank = _rank_within_stream(table)
    return table.with_columns(
        pl.lit(split_name).alias("split_name"),
        pl.lit(SPLIT_NAME_TO_ID[split_name]).cast(pl.Int64).alias("split_id"),
        (rank < warmup_events).alias("is_warmup"),
    ).with_columns((~pl.col("is_warmup")).alias("is_scored"))
assert_temporal_splits_disjoint
assert_temporal_splits_disjoint(*tables: DataFrame) -> None

Raise if any split shares event ids with another split.

Source code in graphids/core/data/preprocessing/temporal.py
def assert_temporal_splits_disjoint(*tables: pl.DataFrame) -> None:
    """Raise if any split shares event ids with another split."""
    seen: set[int] = set()
    for table in tables:
        ids = _event_ids(table)
        overlap = seen & ids
        if overlap:
            sample = sorted(overlap)[:5]
            raise ValueError(f"temporal splits share event_id values: {sample}")
        seen |= ids
build_temporal_event_table
build_temporal_event_table(df: DataFrame, *, id_col: str = 'node_id', raw_id_col: str = 'arb_id', timestamp_col: str = 'timestamp', num_unknown_buckets: int = 128) -> pl.DataFrame

Build one temporal event per row, preserving stream provenance.

The first row in each stream is represented as a self-event (src_id == dst_id) so event counts match normalized row counts.

Source code in graphids/core/data/preprocessing/temporal.py
def build_temporal_event_table(
    df: pl.DataFrame,
    *,
    id_col: str = "node_id",
    raw_id_col: str = "arb_id",
    timestamp_col: str = "timestamp",
    num_unknown_buckets: int = 128,
) -> pl.DataFrame:
    """Build one temporal event per row, preserving stream provenance.

    The first row in each stream is represented as a self-event
    (``src_id == dst_id``) so event counts match normalized row counts.
    """
    if id_col not in df.columns:
        raise ValueError(f"missing mapped id column {id_col!r}")
    if timestamp_col not in df.columns:
        raise ValueError(f"missing timestamp column {timestamp_col!r}")

    df = _ensure_columns(
        df,
        {
            raw_id_col: "",
            "attack": 0,
            "attack_type": 0,
            "entropy": 0.0,
            "vehicle_id": "",
            "source_dir": "",
            "source_file": "",
        },
    )
    df = _ensure_columns(df, {col: 0.0 for col in TEMPORAL_BYTE_COLS})
    stream_keys = ["vehicle_id", "source_dir", "source_file"]
    rows = df.with_row_index("row_index").with_columns(
        pl.col("row_index").cast(pl.Int64),
        pl.col(id_col).cast(pl.Int64).alias("dst_id"),
        pl.col(raw_id_col).cast(pl.Utf8).alias("dst_raw"),
        pl.col(timestamp_col).cast(pl.Float64).alias("timestamp"),
        pl.col("attack").fill_null(0).cast(pl.Int64),
        pl.col("attack_type").fill_null(0).cast(pl.Int64),
        *(pl.col(c).fill_null(0).cast(pl.Float32) for c in TEMPORAL_BYTE_COLS),
        pl.col("entropy").fill_null(0).cast(pl.Float32),
    )
    streams = rows.select(stream_keys).unique(maintain_order=True).with_row_index("stream_id")
    rows = (
        rows.join(streams, on=stream_keys, how="left")
        .with_columns(pl.col("stream_id").cast(pl.Int64))
        .sort(["stream_id", "timestamp", "row_index"])
    )

    delta_exprs = [
        (pl.col(c) - pl.col(c).shift(1).over("stream_id"))
        .fill_null(0)
        .cast(pl.Float32)
        .alias(f"{c}_delta")
        for c in TEMPORAL_BYTE_COLS
    ]
    table = rows.with_columns(
        pl.col("dst_id").shift(1).over("stream_id").fill_null(pl.col("dst_id")).cast(pl.Int64).alias("src_id"),
        pl.col("dst_raw").shift(1).over("stream_id").fill_null(pl.col("dst_raw")).alias("src_raw"),
        (pl.col("timestamp") - pl.col("timestamp").shift(1).over("stream_id"))
        .fill_null(0)
        .cast(pl.Float32)
        .alias("iat"),
        *delta_exprs,
    )
    table = table.with_columns(
        (pl.col("src_id") == 0).alias("src_is_unknown"),
        (pl.col("dst_id") == 0).alias("dst_is_unknown"),
        (
            pl.col("stream_id").shift(-1).is_null()
            | (pl.col("stream_id").shift(-1) != pl.col("stream_id"))
        ).alias("reset_after"),
        (pl.col("attack") > 0).cast(pl.Int64).alias("y"),
    )
    table = _with_unknown_buckets(table, num_unknown_buckets=num_unknown_buckets)
    msg_feature_cols = [
        c
        for c in TEMPORAL_MSG_COL_ORDER
        if c
        not in {
            "src_is_unknown",
            "dst_is_unknown",
            "src_unknown_bucket",
            "dst_unknown_bucket",
        }
    ]
    return table.with_row_index("event_id").with_columns(pl.col("event_id").cast(pl.Int64)).select(
        "event_id",
        "vehicle_id",
        "source_dir",
        "source_file",
        "row_index",
        "timestamp",
        "src_id",
        "dst_id",
        "src_raw",
        "dst_raw",
        "src_is_unknown",
        "dst_is_unknown",
        "src_unknown_bucket",
        "dst_unknown_bucket",
        "stream_id",
        "reset_after",
        *msg_feature_cols,
        "y",
        "attack_type",
    )
mark_split_start_self_events
mark_split_start_self_events(table: DataFrame) -> pl.DataFrame

Remove transition features that cross into the start of a split.

Source code in graphids/core/data/preprocessing/temporal.py
def mark_split_start_self_events(table: pl.DataFrame) -> pl.DataFrame:
    """Remove transition features that cross into the start of a split."""
    if table.is_empty():
        return table
    first_in_stream = _rank_within_stream(table) == 0
    return table.with_columns(
        pl.when(first_in_stream).then(pl.col("dst_id")).otherwise(pl.col("src_id")).alias("src_id"),
        pl.when(first_in_stream).then(pl.col("dst_raw")).otherwise(pl.col("src_raw")).alias("src_raw"),
        pl.when(first_in_stream)
        .then(pl.col("dst_is_unknown"))
        .otherwise(pl.col("src_is_unknown"))
        .alias("src_is_unknown"),
        pl.when(first_in_stream)
        .then(pl.col("dst_unknown_bucket"))
        .otherwise(pl.col("src_unknown_bucket"))
        .alias("src_unknown_bucket"),
        *[
            pl.when(first_in_stream).then(0.0).otherwise(pl.col(c)).cast(pl.Float32).alias(c)
            for c in (*TEMPORAL_DELTA_COLS, "iat")
        ],
    )
mark_terminal_reset
mark_terminal_reset(table: DataFrame) -> pl.DataFrame

Ensure the final event in each sliced stream resets downstream state.

Source code in graphids/core/data/preprocessing/temporal.py
def mark_terminal_reset(table: pl.DataFrame) -> pl.DataFrame:
    """Ensure the final event in each sliced stream resets downstream state."""
    if table.is_empty():
        return table
    if "stream_id" not in table.columns:
        last_event_id = table["event_id"][-1]
        return table.with_columns(
            pl.when(pl.col("event_id") == last_event_id)
            .then(True)
            .otherwise(pl.col("reset_after"))
            .alias("reset_after")
        )
    terminal = (
        table.group_by("stream_id")
        .agg(pl.col("event_id").max().alias("event_id"))
        .with_columns(pl.lit(True).alias("_is_terminal_event"))
    )
    return table.join(terminal, on=["stream_id", "event_id"], how="left").with_columns(
        pl.col("_is_terminal_event").fill_null(False)
    ).with_columns(
        pl.when(pl.col("_is_terminal_event"))
        .then(True)
        .otherwise(pl.col("reset_after"))
        .alias("reset_after")
    ).drop("_is_terminal_event")
prepare_temporal_eval_table
prepare_temporal_eval_table(table: DataFrame, *, split_name: str = 'test', warmup_events: int = 0) -> pl.DataFrame

Prepare validation/test-style full streams with warmup/scoring masks.

Source code in graphids/core/data/preprocessing/temporal.py
def prepare_temporal_eval_table(
    table: pl.DataFrame,
    *,
    split_name: str = "test",
    warmup_events: int = 0,
) -> pl.DataFrame:
    """Prepare validation/test-style full streams with warmup/scoring masks."""
    return add_temporal_split_masks(
        mark_terminal_reset(mark_split_start_self_events(table)),
        split_name=split_name,
        warmup_events=warmup_events,
    )
split_temporal_train_val_tables
split_temporal_train_val_tables(table: DataFrame, *, val_fraction: float, val_warmup_events: int = 0) -> tuple[pl.DataFrame, pl.DataFrame]

Chronologically split each stream into train and validation intervals.

Source code in graphids/core/data/preprocessing/temporal.py
def split_temporal_train_val_tables(
    table: pl.DataFrame,
    *,
    val_fraction: float,
    val_warmup_events: int = 0,
) -> tuple[pl.DataFrame, pl.DataFrame]:
    """Chronologically split each stream into train and validation intervals."""
    train_ids, val_ids = _split_event_ids_by_stream(table, val_fraction=val_fraction)
    train = table.filter(pl.col("event_id").is_in(train_ids))
    val = table.filter(pl.col("event_id").is_in(val_ids))
    train = add_temporal_split_masks(
        mark_terminal_reset(mark_split_start_self_events(train)),
        split_name="train",
        warmup_events=0,
    )
    val = add_temporal_split_masks(
        mark_terminal_reset(mark_split_start_self_events(val)),
        split_name="val",
        warmup_events=val_warmup_events,
    )
    assert_temporal_splits_disjoint(train, val)
    return train, val
temporal_to_pyg
temporal_to_pyg(table: DataFrame) -> TemporalData

Pack a temporal event table into PyG TemporalData tensors.

Source code in graphids/core/data/preprocessing/temporal.py
def temporal_to_pyg(table: pl.DataFrame) -> TemporalData:
    """Pack a temporal event table into PyG ``TemporalData`` tensors."""
    optional_tensors = {}
    optional_specs = {
        "is_warmup": pl.Boolean,
        "is_scored": pl.Boolean,
        "split_id": pl.Int64,
    }
    for col, dtype in optional_specs.items():
        if col in table.columns:
            optional_tensors[col] = _tensor(table, col, dtype=dtype)
    data = TemporalData(
        src=_tensor(table, "src_id", dtype=pl.Int64),
        dst=_tensor(table, "dst_id", dtype=pl.Int64),
        t=_tensor(table, "timestamp", dtype=pl.Float32),
        msg=_tensor(table, list(TEMPORAL_MSG_COL_ORDER), dtype=pl.Float32),
        y=_tensor(table, "y", dtype=pl.Int64),
        attack_type=_tensor(table, "attack_type", dtype=pl.Int64),
        stream_id=_tensor(table, "stream_id", dtype=pl.Int64),
        reset_after=_tensor(table, "reset_after", dtype=pl.Boolean),
        event_id=_tensor(table, "event_id", dtype=pl.Int64),
        **optional_tensors,
    )
    if "split_name" in table.columns:
        names = table["split_name"].unique().to_list()
        # ``TemporalData`` treats normal attributes as event fields and slices
        # them in ``TemporalDataLoader``. Keep this human-readable label as
        # object metadata; the tensor split contract is ``split_id``.
        object.__setattr__(data, "split_name", str(names[0]) if len(names) == 1 else "mixed")
    return data

vocab

Vocabulary scan, digest, persist, and load primitives.

load_vocab
load_vocab(path: Path) -> tuple[dict[str, int], str]

Return (entries, digest) from a persisted vocab file.

Source code in graphids/core/data/preprocessing/vocab.py
def load_vocab(path: Path) -> tuple[dict[str, int], str]:
    """Return ``(entries, digest)`` from a persisted vocab file."""
    payload = json.loads(path.read_text())
    return payload["entries"], payload["digest"]
persist_vocab
persist_vocab(vocab: dict[Any, int], path: Path) -> str

Atomic write and return the digest.

Source code in graphids/core/data/preprocessing/vocab.py
def persist_vocab(vocab: dict[Any, int], path: Path) -> str:
    """Atomic write and return the digest."""
    digest = vocab_digest(vocab)
    payload = {
        "digest": digest,
        "unk_index": UNK_INDEX,
        "entries": {str(k): v for k, v in vocab.items()},
    }
    atomic_write_text(path, json.dumps(payload, indent=2, sort_keys=True))
    return digest
scan_arb_ids
scan_arb_ids(raw_dir: Path, source_dirs: list[str]) -> list[Any]

Sorted unique arb_id across every CSV under source_dirs.

Source code in graphids/core/data/preprocessing/vocab.py
def scan_arb_ids(raw_dir: Path, source_dirs: list[str]) -> list[Any]:
    """Sorted unique ``arb_id`` across every CSV under ``source_dirs``."""
    if not source_dirs:
        raise ValueError("source_dirs is empty; cannot scan for arb_ids")
    frames: list[pl.LazyFrame] = []
    for sub in source_dirs:
        sub_path = raw_dir / sub
        if not sub_path.is_dir():
            raise FileNotFoundError(f"Source dir missing: {sub_path}")
        for csv_path in sorted(sub_path.rglob("*.csv")):
            lf = pl.scan_csv(csv_path)
            cols = lf.collect_schema().names()
            col = "arbitration_id" if "arbitration_id" in cols else "arb_id"
            if col not in cols:
                raise ValueError(
                    f"{csv_path} has neither arbitration_id nor arb_id; got {cols!r}"
                )
            frames.append(lf.select(pl.col(col).alias("arb_id")))
    if not frames:
        raise ValueError(f"No CSVs under {source_dirs!r} in {raw_dir}")
    return pl.concat(frames).collect()["arb_id"].unique().sort().to_list()
vocab_digest
vocab_digest(vocab: dict[Any, int]) -> str

SHA256 over (id, index) pairs sorted by index.

Source code in graphids/core/data/preprocessing/vocab.py
def vocab_digest(vocab: dict[Any, int]) -> str:
    """SHA256 over ``(id, index)`` pairs sorted by index."""
    canon = json.dumps(
        sorted(((str(k), v) for k, v in vocab.items()), key=lambda kv: kv[1]),
        sort_keys=True,
    )
    return hashlib.sha256(canon.encode()).hexdigest()

state

Process-level dataset cache.

DatasetState dataclass

DatasetState(train: Any, val: Any, test: dict[str, Any])

Ready-to-serve train/val/test splits.

clear_cache

clear_cache() -> None

Drop all cached states. Intended for test teardown.

Source code in graphids/core/data/state.py
def clear_cache() -> None:
    """Drop all cached states. Intended for test teardown."""
    _REGISTRY.clear()

get_or_build

get_or_build(dataset: _CacheableDataset) -> DatasetState

Return cached DatasetState for dataset.

Source code in graphids/core/data/state.py
def get_or_build(dataset: _CacheableDataset) -> DatasetState:
    """Return cached ``DatasetState`` for ``dataset``."""
    key = dataset.cache_key
    state = _REGISTRY.get(key)
    if state is None:
        state = dataset.build()
        _REGISTRY[key] = state
    return state