From a4b0b6b981c0c493ed4dee4d22e4dbc57db3f873 Mon Sep 17 00:00:00 2001 From: Andrew Warrington Date: Thu, 1 Oct 2026 10:55:59 -0700 Subject: [PATCH 1/5] Add iTHP model to EasyTPP --- easy_tpp/model/__init__.py | 4 + easy_tpp/model/ithp.py | 713 ++++++++++++++++++++++++ examples/configs/experiment_config.yaml | 39 ++ tests/test_ithp.py | 133 +++++ 4 files changed, 889 insertions(+) create mode 100644 easy_tpp/model/ithp.py create mode 100644 tests/test_ithp.py diff --git a/easy_tpp/model/__init__.py b/easy_tpp/model/__init__.py index 81800bf..fb12d37 100644 --- a/easy_tpp/model/__init__.py +++ b/easy_tpp/model/__init__.py @@ -2,6 +2,7 @@ from easy_tpp.model.attnhp import AttNHP from easy_tpp.model.basemodel import BaseModel, TorchBaseModel from easy_tpp.model.fullynn import FullyNN +from easy_tpp.model.ithp import ITHP from easy_tpp.model.intensity_free import IntensityFree from easy_tpp.model.nhp import NHP from easy_tpp.model.ode_tpp import ODETPP @@ -15,6 +16,7 @@ TorchANHN = ANHN TorchAttNHP = AttNHP TorchFullyNN = FullyNN +TorchITHP = ITHP TorchIntensityFree = IntensityFree TorchNHP = NHP TorchODETPP = ODETPP @@ -28,6 +30,7 @@ 'ANHN', 'AttNHP', 'FullyNN', + 'ITHP', 'IntensityFree', 'NHP', 'ODETPP', @@ -42,6 +45,7 @@ 'TorchANHN', 'TorchAttNHP', 'TorchFullyNN', + 'TorchITHP', 'TorchIntensityFree', 'TorchNHP', 'TorchODETPP', diff --git a/easy_tpp/model/ithp.py b/easy_tpp/model/ithp.py new file mode 100644 index 0000000..d23aa2d --- /dev/null +++ b/easy_tpp/model/ithp.py @@ -0,0 +1,713 @@ +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from easy_tpp.model.basemodel import BaseModel + + +class _ITHPPositionwiseFeedForward(nn.Module): + """Released ITHP feed-forward block with post-layer normalization.""" + + def __init__(self, d_model, d_inner, dropout): + super().__init__() + self.w_1 = nn.Linear(d_model, d_inner) + self.w_2 = nn.Linear(d_inner, d_model) + self.dropout = nn.Dropout(dropout) + self.layer_norm = nn.LayerNorm(d_model, eps=1e-6) + + def forward(self, inputs): + residual = inputs + outputs = self.dropout(F.gelu(self.w_1(inputs))) + outputs = self.dropout(self.w_2(outputs)) + return self.layer_norm(outputs + residual) + + +class _ITHPDynamicValueAttention(nn.Module): + """ITHP attention without learned query or key projections.""" + + def __init__(self, d_model, d_k, d_v, dropout): + super().__init__() + self.scale = math.sqrt(d_k) + self.w_vs = nn.Linear(2 * d_model, d_v, bias=False) + self.fc = nn.Linear(d_v, 2 * d_model) + nn.init.xavier_uniform_(self.fc.weight) + self.dropout = nn.Dropout(dropout) + self.layer_norm = nn.LayerNorm(2 * d_model, eps=1e-6) + + def project_values(self, source_inputs): + return self.w_vs(source_inputs) + + def forward(self, query_inputs, source_inputs, source_values, allowed_mask): + residual = query_inputs + scores = torch.matmul( + query_inputs / self.scale, + source_inputs.transpose(-1, -2), + ) + + has_history = allowed_mask.any(dim=-1, keepdim=True) + masked_scores = scores.masked_fill(~allowed_mask, -torch.inf) + masked_scores = torch.where( + has_history, + masked_scores, + torch.zeros_like(masked_scores), + ) + attention = F.softmax(masked_scores, dim=-1) + attention = attention * allowed_mask.to(attention.dtype) + attention = self.dropout(attention) + + outputs = torch.matmul(attention, source_values) + outputs = self.dropout(self.fc(outputs)) + outputs = self.layer_norm(outputs + residual) + return outputs, attention + + +class _ITHPEncoderLayer(nn.Module): + """Single released ITHP attention and feed-forward layer.""" + + def __init__(self, d_model, d_inner, d_k, d_v, dropout): + super().__init__() + self.self_attention = _ITHPDynamicValueAttention( + d_model=d_model, + d_k=d_k, + d_v=d_v, + dropout=dropout, + ) + self.feed_forward = _ITHPPositionwiseFeedForward( + d_model=d_v, + d_inner=d_inner, + dropout=dropout, + ) + + def forward( + self, + query_inputs, + source_inputs, + source_values, + allowed_mask, + query_mask, + ): + outputs, attention = self.self_attention( + query_inputs=query_inputs, + source_inputs=source_inputs, + source_values=source_values, + allowed_mask=allowed_mask, + ) + outputs = self.feed_forward(outputs) + outputs = outputs * query_mask.unsqueeze(-1).to(outputs.dtype) + return outputs, attention + + +class ITHP(BaseModel): + """Interpretable Transformer Hawkes Process. + + The architecture follows the KDD 2024 authors' public implementation: + https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process + Paper: https://arxiv.org/abs/2405.16059 + Event indexing, padding, likelihood evaluation, and sampling are adapted + to EasyTPP. Attention has one head by construction; ``num_heads`` is not + used. Select integration with ``model_specs.integration_method`` so the + choice survives EasyTPP's config serialization. + + Our reported runs used Adam epsilon 1e-5 and gradient-norm clipping at 1. + EasyTPP's standard runner does not apply those settings, so fresh training + here does not exactly reproduce that training path. Saved checkpoints load + unchanged. + """ + + SUPPORTED_INTEGRATION_METHODS = {"mc", "trapezoid", "fixed_grid"} + + def __init__(self, model_config): + super().__init__(model_config) + + specs = model_config.model_specs or {} + self.d_model = int(model_config.hidden_size) + self.d_inner = int(specs.get("d_inner", 128)) + self.d_k = int(specs.get("d_k", 16)) + self.d_v = int(specs.get("d_v", 2 * self.d_model)) + self.n_layers = int(model_config.num_layers) + self.dropout = float(model_config.dropout_rate) + + self.integration_method = str(specs.get("integration_method", "mc")).lower() + self.grid_step = float(specs.get("grid_step", 0.1)) + self.query_chunk_size = int(specs.get("query_chunk_size", 256)) + self.max_grid_points_per_interval = int( + specs.get("max_grid_points_per_interval", 4096) + ) + + self.use_type_loss = bool(specs.get("use_type_loss", True)) + self.type_loss_weight = float(specs.get("type_loss_weight", 0.5)) + + self._validate_configuration() + if self.integration_method == "mc": + self.use_mc_samples = True + elif self.integration_method == "trapezoid": + self.use_mc_samples = False + + position_vec = torch.tensor( + [ + math.pow(10000.0, 2.0 * (index // 2) / self.d_model) + for index in range(self.d_model) + ], + dtype=torch.float32, + ) + self.register_buffer("position_vec", position_vec) + + self.encoder_layer = _ITHPEncoderLayer( + d_model=self.d_model, + d_inner=self.d_inner, + d_k=self.d_k, + d_v=self.d_v, + dropout=self.dropout, + ) + self.intensity_decoders = nn.ModuleList( + [ + nn.Sequential(nn.Linear(self.d_v, 1), nn.Softplus()) + for _ in range(self.num_event_types) + ] + ) + self.type_predictor = nn.Linear(self.d_v, self.num_event_types) + + self.to(self.device) + + def _validate_configuration(self): + if self.d_model <= 0 or self.d_model % 2: + raise ValueError("ITHP hidden_size must be a positive even integer.") + if self.d_inner <= 0 or self.d_k <= 0 or self.d_v <= 0: + raise ValueError("ITHP d_inner, d_k, and d_v must be positive.") + if self.n_layers != 1: + raise ValueError("Released ITHP supports num_layers == 1.") + if self.d_v != 2 * self.d_model: + raise ValueError("Released ITHP requires d_v == 2 * hidden_size.") + if self.integration_method not in self.SUPPORTED_INTEGRATION_METHODS: + supported = ", ".join(sorted(self.SUPPORTED_INTEGRATION_METHODS)) + raise ValueError( + f"Unsupported ITHP integration_method={self.integration_method!r}; " + f"choose one of {supported}." + ) + if self.loss_integral_num_sample_per_step <= 0: + raise ValueError( + "ITHP loss_integral_num_sample_per_step must be positive." + ) + if ( + self.integration_method == "trapezoid" + and self.loss_integral_num_sample_per_step < 2 + ): + raise ValueError( + "ITHP trapezoid integration requires at least two samples." + ) + if self.grid_step <= 0: + raise ValueError("ITHP grid_step must be positive.") + if self.query_chunk_size <= 0: + raise ValueError("ITHP query_chunk_size must be positive.") + if self.max_grid_points_per_interval <= 0: + raise ValueError( + "ITHP max_grid_points_per_interval must be positive." + ) + if self.type_loss_weight < 0: + raise ValueError("ITHP type_loss_weight cannot be negative.") + + def make_dtime_loss_samples(self, time_delta_seq): + """Draw independent uniform times for MC or fixed times for quadrature.""" + if self.use_mc_samples: + ratios = torch.rand( + (*time_delta_seq.shape, self.loss_integral_num_sample_per_step), + device=self.device, + ) + else: + ratios = torch.linspace( + 0.0, + 1.0, + self.loss_integral_num_sample_per_step, + device=self.device, + )[None, None, :] + return time_delta_seq[..., None] * ratios + + def compute_temporal_embedding(self, times): + """Encode absolute timestamps with the released sinusoidal features.""" + scaled_times = times.to(self.position_vec.dtype).unsqueeze(-1) + scaled_times = scaled_times / self.position_vec + embeddings = torch.empty_like(scaled_times) + embeddings[..., 0::2] = torch.sin(scaled_times[..., 0::2]) + embeddings[..., 1::2] = torch.cos(scaled_times[..., 1::2]) + return embeddings + + def _compose_inputs(self, times, event_types, non_pad_mask): + temporal = self.compute_temporal_embedding(times) + type_embedding = self.layer_type_emb(event_types.long()) + inputs = torch.cat([temporal, type_embedding], dim=-1) + return inputs * non_pad_mask.unsqueeze(-1).to(inputs.dtype) + + @staticmethod + def _build_history_mask(source_mask, history_indices, query_mask): + source_positions = torch.arange( + source_mask.size(1), + device=source_mask.device, + ).view(1, 1, -1) + allowed = source_positions <= history_indices.unsqueeze(-1) + allowed = allowed & source_mask.unsqueeze(1) + return allowed & query_mask.unsqueeze(-1) + + def _prepare_source(self, source_times, source_types, source_mask): + source_inputs = self._compose_inputs( + times=source_times, + event_types=source_types, + non_pad_mask=source_mask, + ) + source_values = self.encoder_layer.self_attention.project_values( + source_inputs + ) + return source_inputs, source_values + + def _encode_query_chunk( + self, + source_inputs, + source_values, + source_mask, + query_times, + query_types, + history_indices, + query_mask, + ): + query_inputs = self._compose_inputs( + times=query_times, + event_types=query_types, + non_pad_mask=query_mask, + ) + allowed_mask = self._build_history_mask( + source_mask=source_mask, + history_indices=history_indices, + query_mask=query_mask, + ) + return self.encoder_layer( + query_inputs=query_inputs, + source_inputs=source_inputs, + source_values=source_values, + allowed_mask=allowed_mask, + query_mask=query_mask, + ) + + def _marked_intensity_chunks( + self, + source_times, + source_types, + source_mask, + query_times, + history_indices, + query_mask, + ): + source_inputs, source_values = self._prepare_source( + source_times=source_times, + source_types=source_types, + source_mask=source_mask, + ) + + candidate_type_embeddings = self.layer_type_emb( + torch.arange( + self.num_event_types, + dtype=torch.long, + device=query_times.device, + ) + ) + decoder_weights = torch.cat( + [decoder[0].weight for decoder in self.intensity_decoders], + dim=0, + ) + decoder_biases = torch.cat( + [decoder[0].bias for decoder in self.intensity_decoders], + dim=0, + ) + num_queries = query_times.size(1) + for start in range(0, num_queries, self.query_chunk_size): + end = min(start + self.query_chunk_size, num_queries) + chunk_times = query_times[:, start:end] + chunk_history = history_indices[:, start:end] + chunk_mask = query_mask[:, start:end] + + batch_size, chunk_size = chunk_times.shape + temporal_embeddings = self.compute_temporal_embedding(chunk_times) + temporal_embeddings = temporal_embeddings.unsqueeze(1).expand( + batch_size, + self.num_event_types, + chunk_size, + self.d_model, + ) + type_embeddings = candidate_type_embeddings.view( + 1, + self.num_event_types, + 1, + self.d_model, + ).expand( + batch_size, + self.num_event_types, + chunk_size, + self.d_model, + ) + query_inputs = torch.cat( + [temporal_embeddings, type_embeddings], + dim=-1, + ).reshape(batch_size, self.num_event_types * chunk_size, -1) + expanded_mask = chunk_mask.unsqueeze(1).expand( + batch_size, + self.num_event_types, + chunk_size, + ).reshape(batch_size, -1) + query_inputs = ( + query_inputs * expanded_mask.unsqueeze(-1).to(query_inputs.dtype) + ) + expanded_history = chunk_history.unsqueeze(1).expand( + batch_size, + self.num_event_types, + chunk_size, + ).reshape(batch_size, -1) + allowed_mask = self._build_history_mask( + source_mask=source_mask, + history_indices=expanded_history, + query_mask=expanded_mask, + ) + enc_output, _ = self.encoder_layer( + query_inputs=query_inputs, + source_inputs=source_inputs, + source_values=source_values, + allowed_mask=allowed_mask, + query_mask=expanded_mask, + ) + enc_output = enc_output.reshape( + batch_size, + self.num_event_types, + chunk_size, + self.d_v, + ) + + logits = torch.einsum( + "btqd,td->btq", + enc_output, + decoder_weights, + ) + logits = logits + decoder_biases.view(1, -1, 1) + chunk_intensities = F.softplus(logits).transpose(1, 2) + chunk_intensities = ( + chunk_intensities + * chunk_mask.unsqueeze(-1).to(chunk_intensities.dtype) + ) + yield start, end, chunk_intensities + + def _compute_marked_intensities( + self, + source_times, + source_types, + source_mask, + query_times, + history_indices, + query_mask, + ): + intensity_chunks = [ + intensities + for _, _, intensities in self._marked_intensity_chunks( + source_times=source_times, + source_types=source_types, + source_mask=source_mask, + query_times=query_times, + history_indices=history_indices, + query_mask=query_mask, + ) + ] + if not intensity_chunks: + return torch.empty( + (*query_times.shape, self.num_event_types), + dtype=self.layer_type_emb.weight.dtype, + device=query_times.device, + ) + return torch.cat(intensity_chunks, dim=1) + + def forward(self, time_seqs, type_seqs, batch_non_pad_mask): + """Compute marked intensities for events after the first event.""" + source_times = time_seqs[:, :-1] + source_types = type_seqs[:, :-1] + source_mask = batch_non_pad_mask[:, :-1].bool() + query_times = time_seqs[:, 1:] + query_mask = batch_non_pad_mask[:, 1:].bool() + history_indices = torch.arange( + query_times.size(1), + device=query_times.device, + ).unsqueeze(0).expand(query_times.size(0), -1) + + return self._compute_marked_intensities( + source_times=source_times, + source_types=source_types, + source_mask=source_mask, + query_times=query_times, + history_indices=history_indices, + query_mask=query_mask, + ) + + def _compute_type_loss(self, time_seqs, type_seqs, batch_non_pad_mask): + source_times = time_seqs[:, :-1] + source_types = type_seqs[:, :-1] + source_mask = batch_non_pad_mask[:, :-1].bool() + query_mask = batch_non_pad_mask[:, 1:].bool() + num_queries = source_times.size(1) + history_indices = torch.arange( + num_queries, + device=source_times.device, + ).unsqueeze(0).expand(source_times.size(0), -1) - 1 + + source_inputs, source_values = self._prepare_source( + source_times=source_times, + source_types=source_types, + source_mask=source_mask, + ) + enc_output, _ = self._encode_query_chunk( + source_inputs=source_inputs, + source_values=source_values, + source_mask=source_mask, + query_times=source_times, + query_types=source_types, + history_indices=history_indices, + query_mask=query_mask, + ) + logits = self.type_predictor(enc_output) + return F.cross_entropy( + logits.transpose(1, 2), + type_seqs[:, 1:], + ignore_index=self.pad_token_id, + reduction="sum", + ) + + def compute_intensities_at_sample_times( + self, + time_seqs, + time_delta_seqs, + type_seqs, + sample_dtimes, + **kwargs, + ): + """Compute marked intensities after each supplied source event.""" + del time_delta_seqs + if sample_dtimes.dim() != 3: + raise ValueError( + "ITHP sample_dtimes must have shape [batch, sequence, samples]." + ) + + source_mask = kwargs.get("source_mask") + if source_mask is None: + source_mask = type_seqs.ne(self.pad_token_id) + source_mask = source_mask.bool() + + interval_mask = kwargs.get("query_mask") + if interval_mask is None: + interval_mask = source_mask + interval_mask = interval_mask.bool() + + batch_size, seq_len, num_samples = sample_dtimes.shape + query_times = time_seqs.unsqueeze(-1) + sample_dtimes + query_times = query_times.reshape(batch_size, seq_len * num_samples) + history_indices = torch.arange( + seq_len, + device=time_seqs.device, + ).view(1, seq_len, 1).expand(batch_size, seq_len, num_samples) + history_indices = history_indices.reshape(batch_size, -1) + query_mask = interval_mask.unsqueeze(-1).expand( + batch_size, + seq_len, + num_samples, + ).reshape(batch_size, -1) + + intensities = self._compute_marked_intensities( + source_times=time_seqs, + source_types=type_seqs, + source_mask=source_mask, + query_times=query_times, + history_indices=history_indices, + query_mask=query_mask, + ) + intensities = intensities.reshape( + batch_size, + seq_len, + num_samples, + self.num_event_types, + ) + + if kwargs.get("compute_last_step_only", False): + return intensities[:, -1:, :, :] + return intensities + + def _integrate_fixed_grid( + self, + source_times, + source_types, + source_mask, + time_delta_seqs, + interval_mask, + ): + cell_counts = torch.ceil(time_delta_seqs / self.grid_step).long() + cell_counts = torch.where( + interval_mask & time_delta_seqs.gt(0), + cell_counts, + torch.zeros_like(cell_counts), + ) + max_cells = int(cell_counts.max().item()) if cell_counts.numel() else 0 + if max_cells > self.max_grid_points_per_interval: + raise ValueError( + "ITHP fixed_grid requires " + f"{max_cells} points in one interval, exceeding " + f"max_grid_points_per_interval=" + f"{self.max_grid_points_per_interval}. Use integration_method=" + "'mc' or increase grid_step." + ) + if max_cells == 0: + return torch.zeros_like(time_delta_seqs) + + batch_size, num_intervals = time_delta_seqs.shape + cell_index = torch.arange( + max_cells, + device=time_delta_seqs.device, + dtype=time_delta_seqs.dtype, + ).view(1, 1, -1) + cell_start = cell_index * self.grid_step + cell_width = torch.clamp( + time_delta_seqs.unsqueeze(-1) - cell_start, + min=0.0, + max=self.grid_step, + ) + cell_mask = interval_mask.unsqueeze(-1) & cell_width.gt(0) + query_times = ( + source_times.unsqueeze(-1) + + cell_start + + 0.5 * cell_width + ) + + flat_query_times = query_times.reshape(batch_size, -1) + flat_query_mask = cell_mask.reshape(batch_size, -1) + flat_weights = cell_width.reshape(batch_size, -1) + flat_history = torch.arange( + num_intervals, + device=time_delta_seqs.device, + ).view(1, num_intervals, 1).expand( + batch_size, + num_intervals, + max_cells, + ).reshape(batch_size, -1) + flat_intervals = torch.arange( + num_intervals, + device=time_delta_seqs.device, + ).view(1, num_intervals, 1).expand( + batch_size, + num_intervals, + max_cells, + ).reshape(batch_size, -1) + + non_event_ll = torch.zeros_like(time_delta_seqs) + for start, end, intensities in self._marked_intensity_chunks( + source_times=source_times, + source_types=source_types, + source_mask=source_mask, + query_times=flat_query_times, + history_indices=flat_history, + query_mask=flat_query_mask, + ): + weighted_intensity = ( + intensities.sum(dim=-1) * flat_weights[:, start:end] + ) + interval_indices = flat_intervals[:, start:end] + contribution = torch.zeros_like(non_event_ll).scatter_add( + dim=1, + index=interval_indices, + src=weighted_intensity, + ) + non_event_ll = non_event_ll + contribution + return non_event_ll + + def _compute_grid_loglikelihood( + self, + time_delta_seqs, + lambda_at_event, + non_event_ll, + seq_mask, + type_seqs, + ): + lambda_at_event = lambda_at_event + self.eps + log_marked_event_lambdas = lambda_at_event.log() + event_ll = -F.nll_loss( + log_marked_event_lambdas.permute(0, 2, 1), + target=type_seqs, + ignore_index=self.pad_token_id, + reduction="none", + ) + non_event_ll = non_event_ll * seq_mask * time_delta_seqs.gt(0) + num_events = int(seq_mask.sum().item()) + return event_ll, non_event_ll, num_events + + def loglike_loss(self, batch=None, **kwargs): + """Compute ITHP likelihood and optional auxiliary type loss.""" + ( + time_seqs, + time_delta_seqs, + type_seqs, + batch_non_pad_mask, + _, + ) = self.resolve_batch_inputs(batch, kwargs) + if time_seqs.size(1) < 2: + raise ValueError("ITHP requires sequences with at least two events.") + + target_mask = batch_non_pad_mask[:, 1:].bool() + lambda_at_event = self.forward( + time_seqs=time_seqs, + type_seqs=type_seqs, + batch_non_pad_mask=batch_non_pad_mask, + ) + + source_times = time_seqs[:, :-1] + source_types = type_seqs[:, :-1] + source_mask = batch_non_pad_mask[:, :-1].bool() + target_dtimes = time_delta_seqs[:, 1:] + target_types = type_seqs[:, 1:] + + if self.integration_method == "fixed_grid": + non_event_ll = self._integrate_fixed_grid( + source_times=source_times, + source_types=source_types, + source_mask=source_mask, + time_delta_seqs=target_dtimes, + interval_mask=target_mask, + ) + event_ll, non_event_ll, num_events = self._compute_grid_loglikelihood( + time_delta_seqs=target_dtimes, + lambda_at_event=lambda_at_event, + non_event_ll=non_event_ll, + seq_mask=target_mask, + type_seqs=target_types, + ) + else: + sample_dtimes = self.make_dtime_loss_samples(target_dtimes) + lambda_t_sample = self.compute_intensities_at_sample_times( + time_seqs=source_times, + time_delta_seqs=time_delta_seqs[:, :-1], + type_seqs=source_types, + sample_dtimes=sample_dtimes, + source_mask=source_mask, + query_mask=target_mask, + ) + (event_ll, non_event_ll, num_events) = self.compute_loglikelihood( + time_delta_seq=target_dtimes, + lambda_at_event=lambda_at_event, + lambdas_loss_samples=lambda_t_sample, + seq_mask=target_mask, + type_seq=target_types, + ) + + if num_events == 0: + raise ValueError("ITHP batch contains no target events.") + + nll = -(event_ll - non_event_ll).sum() + loss = nll + if self.training and self.use_type_loss and self.type_loss_weight > 0: + type_loss = self._compute_type_loss( + time_seqs=time_seqs, + type_seqs=type_seqs, + batch_non_pad_mask=batch_non_pad_mask, + ) + loss = loss + self.type_loss_weight * type_loss + + return loss, num_events diff --git a/examples/configs/experiment_config.yaml b/examples/configs/experiment_config.yaml index d506043..795e11f 100644 --- a/examples/configs/experiment_config.yaml +++ b/examples/configs/experiment_config.yaml @@ -694,6 +694,45 @@ S2P2_train: relative_time: True # If True, predicts the scaling factor to be applied to the dynamics between each pair of subsequent events. See Sec. 3.3 of the paper. +# ITHP: Interpretable Transformer Hawkes Process (KDD 2024) +# Architecture adapted from https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process +# This Taobao example uses the selected width, LR, and epoch budget. The +# reported runs also used Adam epsilon 1e-5 and gradient-norm clipping at 1; +# stock EasyTPP does not apply those settings. +ITHP_train: + base_config: + stage: train + backend: torch + dataset_id: taobao + runner_id: std_tpp + model_id: ITHP + base_dir: './checkpoints/' + trainer_config: + batch_size: 16 + max_epoch: 300 + shuffle: False + optimizer: adam + learning_rate: 5.e-4 + valid_freq: 1 + use_tfb: False + metrics: null + seed: 2019 + gpu: -1 + model_config: + hidden_size: 16 + num_layers: 1 + dropout_rate: 0.0 + loss_integral_num_sample_per_step: 20 + model_specs: + d_inner: 128 + d_k: 16 + d_v: 32 + integration_method: mc + query_chunk_size: 256 + use_type_loss: true + type_loss_weight: 0.5 + + # WSM-THP: Weighted Score Matching Temporal Hawkes Process # Ref: "Is Score Matching Suitable for Estimating Point Processes?" NeurIPS 2024 # https://arxiv.org/abs/2512.04617 diff --git a/tests/test_ithp.py b/tests/test_ithp.py new file mode 100644 index 0000000..0fba6f2 --- /dev/null +++ b/tests/test_ithp.py @@ -0,0 +1,133 @@ +import torch + +from easy_tpp.config_factory.model_config import ModelConfig +from easy_tpp.model import BaseModel, ITHP + + +def make_config(integration_method='mc'): + return ModelConfig( + model_id='ITHP', + hidden_size=4, + num_layers=1, + num_event_types=3, + num_event_types_pad=4, + event_pad_index=3, + gpu=-1, + training=True, + loss_integral_num_sample_per_step=5, + model_specs={ + 'd_inner': 12, + 'd_k': 4, + 'd_v': 8, + 'integration_method': integration_method, + 'grid_step': 0.2, + 'query_chunk_size': 5, + 'use_type_loss': True, + 'type_loss_weight': 0.5, + }, + ) + + +def make_batch(padded=False): + times = torch.tensor([[0.0, 0.4, 1.0]]) + dtimes = torch.tensor([[0.0, 0.4, 0.6]]) + types = torch.tensor([[0, 1, 2]]) + mask = torch.tensor([[True, True, True]]) + if padded: + times = torch.cat((times, times[:, -1:].expand(-1, 2)), dim=1) + dtimes = torch.cat((dtimes, torch.zeros(1, 2)), dim=1) + types = torch.cat((types, torch.full((1, 2), 3)), dim=1) + mask = torch.cat((mask, torch.zeros(1, 2, dtype=torch.bool)), dim=1) + return { + 'time_seqs': times, + 'time_delta_seqs': dtimes, + 'type_seqs': types, + 'seq_non_pad_mask': mask, + 'attention_mask': torch.zeros((1, times.size(1), times.size(1)), dtype=torch.bool), + } + + +def test_registration_named_batch_and_gradients(): + model = BaseModel.generate_model_from_config(make_config()) + assert isinstance(model, ITHP) + torch.manual_seed(7) + loss, count = model.loglike_loss(**make_batch()) + assert count == 2 + assert torch.isfinite(loss) + loss.backward() + assert model.encoder_layer.self_attention.w_vs.weight.grad is not None + assert torch.isfinite(model.encoder_layer.self_attention.w_vs.weight.grad).all() + + +def test_mc_is_random_and_trapezoid_is_fixed(): + dtimes = torch.ones(1, 2) + mc_model = ITHP(make_config('mc')) + first = mc_model.make_dtime_loss_samples(dtimes) + second = mc_model.make_dtime_loss_samples(dtimes) + assert not torch.equal(first, second) + assert ((first >= 0) & (first <= 1)).all() + + trapezoid_model = ITHP(make_config('trapezoid')) + expected = torch.linspace(0, 1, 5).expand(1, 2, -1) + assert torch.equal(trapezoid_model.make_dtime_loss_samples(dtimes), expected) + + +def test_padding_does_not_change_valid_likelihood(): + model = ITHP(make_config()) + model.eval() + torch.manual_seed(17) + plain_loss, plain_count = model.loglike_loss(**make_batch()) + torch.manual_seed(17) + padded_loss, padded_count = model.loglike_loss(**make_batch(padded=True)) + assert plain_count == padded_count == 2 + torch.testing.assert_close(plain_loss, padded_loss, rtol=0, atol=1e-5) + + +def test_future_event_does_not_change_earlier_intensity(): + model = ITHP(make_config()) + model.eval() + batch = make_batch() + baseline = model(batch['time_seqs'], batch['type_seqs'], batch['seq_non_pad_mask']) + changed_times = batch['time_seqs'].clone() + changed_types = batch['type_seqs'].clone() + changed_times[:, -1] = 9.0 + changed_types[:, -1] = 0 + changed = model(changed_times, changed_types, batch['seq_non_pad_mask']) + torch.testing.assert_close(baseline[:, 0], changed[:, 0]) + + +def test_auxiliary_type_loss_applies_only_during_training(): + model = ITHP(make_config('fixed_grid')) + model.train() + train_loss, _ = model.loglike_loss(**make_batch()) + model.eval() + valid_loss, _ = model.loglike_loss(**make_batch()) + assert train_loss > valid_loss + + +def test_integration_modes_and_intensity_shape(): + batch = make_batch() + for method in ('mc', 'trapezoid', 'fixed_grid'): + model = ITHP(make_config(method)) + model.eval() + loss, count = model.loglike_loss(**batch) + assert count == 2 + assert torch.isfinite(loss) + intensities = model.compute_intensities_at_sample_times( + time_seqs=batch['time_seqs'][:, :-1], + time_delta_seqs=batch['time_delta_seqs'][:, :-1], + type_seqs=batch['type_seqs'][:, :-1], + sample_dtimes=torch.full((1, 2, 3), 0.1), + ) + assert intensities.shape == (1, 2, 3, 3) + assert (intensities > 0).all() + + +def test_generated_config_reconstructs_checkpoint_shape(): + config = make_config('trapezoid') + original = ITHP(config) + restored_config = ModelConfig.parse_from_yaml_config(config.get_yaml_config()) + restored = ITHP(restored_config) + restored.load_state_dict(original.state_dict(), strict=True) + assert restored.integration_method == 'trapezoid' + assert restored.use_mc_samples is False From ce4c7f7472f8fcd1e0e308502ad924e4bb022f24 Mon Sep 17 00:00:00 2001 From: Andrew Warrington Date: Thu, 1 Oct 2026 11:10:59 -0700 Subject: [PATCH 2/5] Document iTHP source and integration choices --- easy_tpp/model/ithp.py | 68 +++++++++++++++++++++---- examples/configs/experiment_config.yaml | 2 + tests/test_ithp.py | 20 ++++++++ 3 files changed, 81 insertions(+), 9 deletions(-) diff --git a/easy_tpp/model/ithp.py b/easy_tpp/model/ithp.py index d23aa2d..0640e0a 100644 --- a/easy_tpp/model/ithp.py +++ b/easy_tpp/model/ithp.py @@ -25,10 +25,12 @@ def forward(self, inputs): class _ITHPDynamicValueAttention(nn.Module): - """ITHP attention without learned query or key projections.""" + """Public-code attention: direct query/key dot product, one attention map.""" def __init__(self, d_model, d_k, d_v, dropout): super().__init__() + # Released SubLayers.dynamic_v_attention scales by sqrt(d_k), whereas + # paper Eqs. (4)-(5) use sqrt(2M) for the concatenated query/key. self.scale = math.sqrt(d_k) self.w_vs = nn.Linear(2 * d_model, d_v, bias=False) self.fc = nn.Linear(d_v, 2 * d_model) @@ -100,20 +102,64 @@ def forward( class ITHP(BaseModel): - """Interpretable Transformer Hawkes Process. + """Interpretable Transformer Hawkes Process, public-code architecture. - The architecture follows the KDD 2024 authors' public implementation: + Paper: Meng et al., KDD 2024, https://arxiv.org/abs/2405.16059 + Authors' code, pinned at commit 5db1bb78f3323667e2cef478e177cd35971c4b43: https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process - Paper: https://arxiv.org/abs/2405.16059 - Event indexing, padding, likelihood evaluation, and sampling are adapted - to EasyTPP. Attention has one head by construction; ``num_heads`` is not - used. Select integration with ``model_specs.integration_method`` so the - choice survives EasyTPP's config serialization. + + Source-to-implementation map: + + * Paper Sec. 4.1/Eq. (3) concatenates sinusoidal time and learned type + embeddings; Sec. 4.2/Eqs. (4)-(5) use unprojected queries and keys. + See public ``transformer/Models.py`` lines 50-106 and + ``transformer/SubLayers.py`` lines 194-236. This class retains both. + * Public ``Models.Encoder`` passes ``n_head`` to ``EncoderLayer``, but its + active ``dynamic_v_attention`` never reads that argument. The separate + ``MultiHeadAttention`` class is not called by ``EncoderLayer``. See + ``Models.py`` lines 50-68, ``Layers.py`` lines 9-23, and + ``SubLayers.py`` lines 13-69 and 194-236. ``Main.py`` line 258 defaults + to one head. This model has one attention map by construction; EasyTPP's + generic ``num_heads`` setting (default 2) has no effect here. + * Paper Eq. (7) decodes the weighted value sum directly and scales its + dot product by sqrt(2M). The public attention instead scales by + sqrt(d_k), then applies a projection, residual, layer norm, and feed- + forward block before its type-specific softplus decoder. See public + ``SubLayers.py`` lines 194-301, ``Layers.py`` lines 9-23, and + ``Models.py`` lines 203-236. We follow that executed decoder, so its + attention weights are not an exact additive Hawkes-kernel decomposition. + * Paper Sec. 4.5/Eq. (8) describes numerical integration on a fine grid. + Public ``Utils.gen_Xne`` lines 104-142 builds a batch-wide 0.1 grid; + ``Main.py`` lines 62-72 passes it to ``Models.Transformer.forward``, + which sums its intensities (lines 203-236). Here the loss is + conditional on the first event and integrates each observed interval: + ``integration_method='mc'`` draws independent uniform times (default), + ``'trapezoid'`` uses evenly spaced nodes including endpoints, and + ``'fixed_grid'`` uses per-interval midpoint cells. This matches the + random MC path used by our EasyTPP experiments, not the authors' + active 0.1-grid script. Public ``Utils.py`` lines 42-59 contains a + separate MC helper, but the active ``Main.py`` path does not call it. + * Paper Sec. 4.5 specifies MLE; public ``Main.py`` lines 62-115 adds + 0.5 times next-type cross entropy during training. We keep that term + during training only; validation/test loss is point-process NLL. + + Event indexing and padding follow EasyTPP. We score events after the first + and mask all-history-empty queries, avoiding first-event future leakage. + Select the integration rule through ``model_specs.integration_method``; + EasyTPP serializes that field but omits ``use_mc_samples`` and ``num_heads``. Our reported runs used Adam epsilon 1e-5 and gradient-norm clipping at 1. EasyTPP's standard runner does not apply those settings, so fresh training here does not exactly reproduce that training path. Saved checkpoints load unchanged. + + Pinned code for the line references above: + https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process/blob/5db1bb78f3323667e2cef478e177cd35971c4b43/transformer/Models.py + https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process/blob/5db1bb78f3323667e2cef478e177cd35971c4b43/transformer/Layers.py + https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process/blob/5db1bb78f3323667e2cef478e177cd35971c4b43/transformer/SubLayers.py + https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process/blob/5db1bb78f3323667e2cef478e177cd35971c4b43/Main.py + https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process/blob/5db1bb78f3323667e2cef478e177cd35971c4b43/Utils.py + Our prior MC sampler: https://github.com/andrewwarrington/HHP/blob/76fbcf00c33f3d2d8e937c5984e9a1e17b218cde/EasyTPP/easy_tpp/model/torch_model/torch_basemodel.py#L170-L197 """ SUPPORTED_INTEGRATION_METHODS = {"mc", "trapezoid", "fixed_grid"} @@ -209,7 +255,11 @@ def _validate_configuration(self): raise ValueError("ITHP type_loss_weight cannot be negative.") def make_dtime_loss_samples(self, time_delta_seq): - """Draw independent uniform times for MC or fixed times for quadrature.""" + """Draw random MC times or fixed trapezoid nodes per interval. + + Upstream EasyTPP's BaseModel currently returns ``linspace`` for both + modes; overriding it preserves the stochastic rule of our runs. + """ if self.use_mc_samples: ratios = torch.rand( (*time_delta_seq.shape, self.loss_integral_num_sample_per_step), diff --git a/examples/configs/experiment_config.yaml b/examples/configs/experiment_config.yaml index 795e11f..59ed33d 100644 --- a/examples/configs/experiment_config.yaml +++ b/examples/configs/experiment_config.yaml @@ -721,6 +721,8 @@ ITHP_train: model_config: hidden_size: 16 num_layers: 1 + # The public CLI defaults to n_head=1, but its active attention ignores it. + # ITHP has one attention map and does not use EasyTPP's num_heads setting. dropout_rate: 0.0 loss_integral_num_sample_per_step: 20 model_specs: diff --git a/tests/test_ithp.py b/tests/test_ithp.py index 0fba6f2..8af638b 100644 --- a/tests/test_ithp.py +++ b/tests/test_ithp.py @@ -72,6 +72,26 @@ def test_mc_is_random_and_trapezoid_is_fixed(): assert torch.equal(trapezoid_model.make_dtime_loss_samples(dtimes), expected) +def test_legacy_num_heads_setting_does_not_change_active_attention(): + one_head_config = make_config('fixed_grid') + one_head_config.num_heads = 1 + four_head_config = make_config('fixed_grid') + four_head_config.num_heads = 4 + + torch.manual_seed(13) + one_head_model = ITHP(one_head_config).eval() + torch.manual_seed(13) + four_head_model = ITHP(four_head_config).eval() + + for key, value in one_head_model.state_dict().items(): + torch.testing.assert_close(value, four_head_model.state_dict()[key]) + batch = make_batch() + torch.testing.assert_close( + one_head_model(batch['time_seqs'], batch['type_seqs'], batch['seq_non_pad_mask']), + four_head_model(batch['time_seqs'], batch['type_seqs'], batch['seq_non_pad_mask']), + ) + + def test_padding_does_not_change_valid_likelihood(): model = ITHP(make_config()) model.eval() From 3a0d2b727300c15096bac3bb1e2ae3c244178e88 Mon Sep 17 00:00:00 2001 From: Andrew Warrington Date: Thu, 1 Oct 2026 11:14:50 -0700 Subject: [PATCH 3/5] Clarify first-event masking provenance --- easy_tpp/model/ithp.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/easy_tpp/model/ithp.py b/easy_tpp/model/ithp.py index 0640e0a..0f5cc8a 100644 --- a/easy_tpp/model/ithp.py +++ b/easy_tpp/model/ithp.py @@ -137,14 +137,18 @@ class ITHP(BaseModel): ``'trapezoid'`` uses evenly spaced nodes including endpoints, and ``'fixed_grid'`` uses per-interval midpoint cells. This matches the random MC path used by our EasyTPP experiments, not the authors' - active 0.1-grid script. Public ``Utils.py`` lines 42-59 contains a + active 0.1-grid script. Public ``Utils.py`` lines 42-59 contain a separate MC helper, but the active ``Main.py`` path does not call it. * Paper Sec. 4.5 specifies MLE; public ``Main.py`` lines 62-115 adds 0.5 times next-type cross entropy during training. We keep that term during training only; validation/test loss is point-process NLL. + * Public ``Models.get_subsequent_mask`` (lines 28-36) masks the diagonal + and all later keys; ``dynamic_v_attention`` (``SubLayers.py`` lines + 222-227) softmaxes the fully masked first row. + Here an empty-history row has zero attention, and EasyTPP's likelihood + scores events only after the first. - Event indexing and padding follow EasyTPP. We score events after the first - and mask all-history-empty queries, avoiding first-event future leakage. + Event indexing and padding follow EasyTPP. Select the integration rule through ``model_specs.integration_method``; EasyTPP serializes that field but omits ``use_mc_samples`` and ``num_heads``. From 32e6eb0316a6247e4e75c4e202f1855b71a4788b Mon Sep 17 00:00:00 2001 From: Andrew Warrington Date: Thu, 1 Oct 2026 11:36:03 -0700 Subject: [PATCH 4/5] Document ITHP paper and public code decisions --- easy_tpp/model/ithp.py | 110 ++++++++++++++++-------- examples/configs/experiment_config.yaml | 2 + 2 files changed, 74 insertions(+), 38 deletions(-) diff --git a/easy_tpp/model/ithp.py b/easy_tpp/model/ithp.py index 0f5cc8a..26fcb39 100644 --- a/easy_tpp/model/ithp.py +++ b/easy_tpp/model/ithp.py @@ -105,50 +105,78 @@ class ITHP(BaseModel): """Interpretable Transformer Hawkes Process, public-code architecture. Paper: Meng et al., KDD 2024, https://arxiv.org/abs/2405.16059 + Relevant pages: https://arxiv.org/pdf/2405.16059#page=3 (Sec. 4.1), + https://arxiv.org/pdf/2405.16059#page=4 (Secs. 4.2-4.3), + https://arxiv.org/pdf/2405.16059#page=5 (Secs. 4.4-4.5), and + https://arxiv.org/pdf/2405.16059#page=9 (Sec. 5.4). Authors' code, pinned at commit 5db1bb78f3323667e2cef478e177cd35971c4b43: https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process - Source-to-implementation map: + Paper-to-code decisions (line references are to the pinned source below): - * Paper Sec. 4.1/Eq. (3) concatenates sinusoidal time and learned type - embeddings; Sec. 4.2/Eqs. (4)-(5) use unprojected queries and keys. - See public ``transformer/Models.py`` lines 50-106 and - ``transformer/SubLayers.py`` lines 194-236. This class retains both. - * Public ``Models.Encoder`` passes ``n_head`` to ``EncoderLayer``, but its - active ``dynamic_v_attention`` never reads that argument. The separate - ``MultiHeadAttention`` class is not called by ``EncoderLayer``. See + * Retained from paper Sec. 4.1/Eq. (3): concatenate sinusoidal time and + learned type embeddings. Sec. 4.2/Eqs. (4)-(5): unprojected queries and + keys with projected values. Public ``Models.py`` lines 50-106 and + ``SubLayers.py`` lines 194-236 implement these choices too. + * The paper's Eqs. (4)-(5) have one attention matrix and do not specify + a head-count hyperparameter. Public ``Models.Encoder`` passes ``n_head`` + to ``EncoderLayer``, but the active ``dynamic_v_attention`` never uses + it; the separate ``MultiHeadAttention`` class is inactive. See public ``Models.py`` lines 50-68, ``Layers.py`` lines 9-23, and ``SubLayers.py`` lines 13-69 and 194-236. ``Main.py`` line 258 defaults - to one head. This model has one attention map by construction; EasyTPP's - generic ``num_heads`` setting (default 2) has no effect here. - * Paper Eq. (7) decodes the weighted value sum directly and scales its - dot product by sqrt(2M). The public attention instead scales by - sqrt(d_k), then applies a projection, residual, layer norm, and feed- - forward block before its type-specific softplus decoder. See public - ``SubLayers.py`` lines 194-301, ``Layers.py`` lines 9-23, and - ``Models.py`` lines 203-236. We follow that executed decoder, so its - attention weights are not an exact additive Hawkes-kernel decomposition. - * Paper Sec. 4.5/Eq. (8) describes numerical integration on a fine grid. - Public ``Utils.gen_Xne`` lines 104-142 builds a batch-wide 0.1 grid; - ``Main.py`` lines 62-72 passes it to ``Models.Transformer.forward``, - which sums its intensities (lines 203-236). Here the loss is - conditional on the first event and integrates each observed interval: - ``integration_method='mc'`` draws independent uniform times (default), - ``'trapezoid'`` uses evenly spaced nodes including endpoints, and - ``'fixed_grid'`` uses per-interval midpoint cells. This matches the - random MC path used by our EasyTPP experiments, not the authors' - active 0.1-grid script. Public ``Utils.py`` lines 42-59 contain a - separate MC helper, but the active ``Main.py`` path does not call it. - * Paper Sec. 4.5 specifies MLE; public ``Main.py`` lines 62-115 adds - 0.5 times next-type cross entropy during training. We keep that term - during training only; validation/test loss is point-process NLL. - * Public ``Models.get_subsequent_mask`` (lines 28-36) masks the diagonal - and all later keys; ``dynamic_v_attention`` (``SubLayers.py`` lines - 222-227) softmaxes the fully masked first row. - Here an empty-history row has zero attention, and EasyTPP's likelihood - scores events only after the first. - - Event indexing and padding follow EasyTPP. + to one head. We therefore have one attention map and ignore EasyTPP's + generic ``num_heads`` setting (whose default is 2). + * Paper Eq. (5) scales the 2M-wide query/key dot product by sqrt(2M). + Public ``dynamic_v_attention`` instead divides by sqrt(d_k) + (``SubLayers.py`` line 222). We retain the executed public rule so the + released architecture and our saved checkpoints have the same scores. + * Paper Eq. (7) writes a type-specific softplus head directly on the + attention-weighted value sum. However, paper Sec. 5.4 explicitly says + the encoder keeps a skip connection and requires M_V=2M; Eq. (7) does + not show how that connection enters the intensity. Public code resolves + this ambiguity with a projection, query residual, layer norm, and GELU + feed-forward block before the type-specific decoder (``SubLayers.py`` + lines 194-301; ``Layers.py`` lines 9-23; ``Models.py`` lines 203-236). + We use that executed path, including its forced single encoder layer + (``Models.py`` lines 65-68), rather than silently inventing an Eq. (7) + variant. Its intensity does not have Eq. (7)'s exact additive form. + * Paper Sec. 4.4/Eq. (7) defines a marked intensity at event and + non-event times using only preceding events as keys. Public + ``Models.Transformer.forward`` loops over candidate types (lines + 203-236) and excludes grid points as keys (``Models.py`` lines 83-96). + We vectorize the same candidate-type computation, using + EasyTPP's zero-based marks instead of the public code's one-based marks. + * Paper Sec. 4.5/Eq. (8) scores events from the beginning of each + sequence and integrates on [0,T] using an unspecified sufficiently + fine grid. Public ``Utils.gen_Xne`` (lines 104-142) instead makes a + batch-wide 0.1 grid and shifts time by the batch minimum; ``Main.py`` + lines 62-72 passes it to ``Models.Transformer.forward``. We condition + on the first observed event, as EasyTPP baselines do, and integrate + each sequence's observed intervals separately to remove dependence + on unrelated batch members. No batch-wide time shift is applied. + * For that integral, ``integration_method='mc'`` draws independent + uniform times per interval by default; ``'trapezoid'`` uses evenly + spaced nodes and ``'fixed_grid'`` uses per-interval midpoint cells. + This replaces the authors' active 0.1 grid, but retains the stochastic + sampler of our original EasyTPP runs. The separate public MC helper + (``Utils.py`` lines 42-59) is not called by the active ``Main.py`` path. + Our reported runs used 20 draws per interval for training/validation + and 100 for test likelihood; sample counts are an approximation choice, + not a number specified by paper Sec. 4.5. Even ``fixed_grid`` with + ``grid_step=0.1`` is a per-interval midpoint rule, not a reproduction + of the public batch-wide grid. + * Paper Eq. (8) describes maximum likelihood alone. Public ``Main.py`` + lines 62-115 adds 0.5 times next-type cross entropy in training. We + retain that public training term but report pure point-process NLL on + validation/test. Public ``Main.py`` lines 32-47 and 177-208 select a + checkpoint using test LL; EasyTPP instead selects by validation LL, + keeping the test split held out. + * Paper Eqs. (5) and (7) require strictly preceding events. Public + ``Models.get_subsequent_mask`` (lines 28-36) masks the entire first + row, then ``dynamic_v_attention`` softmaxes its finite mask values + (``SubLayers.py`` lines 222-227). We zero empty-history attention and + do not score the first event, avoiding future leakage at that row. + Select the integration rule through ``model_specs.integration_method``; EasyTPP serializes that field but omits ``use_mc_samples`` and ``num_heads``. @@ -227,8 +255,10 @@ def _validate_configuration(self): if self.d_inner <= 0 or self.d_k <= 0 or self.d_v <= 0: raise ValueError("ITHP d_inner, d_k, and d_v must be positive.") if self.n_layers != 1: + # Public Models.Encoder forces one layer, despite its n_layers arg. raise ValueError("Released ITHP supports num_layers == 1.") if self.d_v != 2 * self.d_model: + # Paper Sec. 5.4 states that the skip connection requires M_V=2M. raise ValueError("Released ITHP requires d_v == 2 * hidden_size.") if self.integration_method not in self.SUPPORTED_INTEGRATION_METHODS: supported = ", ".join(sorted(self.SUPPORTED_INTEGRATION_METHODS)) @@ -295,6 +325,8 @@ def _compose_inputs(self, times, event_types, non_pad_mask): @staticmethod def _build_history_mask(source_mask, history_indices, query_mask): + # Paper Eqs. (5)/(7) use only prior events; zero-history rows are + # handled in _ITHPDynamicValueAttention instead of softmaxing -1e9. source_positions = torch.arange( source_mask.size(1), device=source_mask.device, @@ -705,6 +737,8 @@ def loglike_loss(self, batch=None, **kwargs): if time_seqs.size(1) < 2: raise ValueError("ITHP requires sequences with at least two events.") + # EasyTPP compares conditional LL after the first observed event. + # Paper Eq. (8) instead writes the full [0,T] sequence likelihood. target_mask = batch_non_pad_mask[:, 1:].bool() lambda_at_event = self.forward( time_seqs=time_seqs, diff --git a/examples/configs/experiment_config.yaml b/examples/configs/experiment_config.yaml index 59ed33d..6222410 100644 --- a/examples/configs/experiment_config.yaml +++ b/examples/configs/experiment_config.yaml @@ -699,6 +699,8 @@ S2P2_train: # This Taobao example uses the selected width, LR, and epoch budget. The # reported runs also used Adam epsilon 1e-5 and gradient-norm clipping at 1; # stock EasyTPP does not apply those settings. +# The reported runs used 20 random samples/interval for training and +# validation, then 100 samples/interval for held-out test log likelihood. ITHP_train: base_config: stage: train From 1d4d72b852579352bf5309dd52a721c7cee933a2 Mon Sep 17 00:00:00 2001 From: Andrew Warrington Date: Thu, 1 Oct 2026 12:05:22 -0700 Subject: [PATCH 5/5] Keep ITHP provenance focused on paper and public code --- easy_tpp/model/ithp.py | 16 +++++++--------- 1 file changed, 7 insertions(+), 9 deletions(-) diff --git a/easy_tpp/model/ithp.py b/easy_tpp/model/ithp.py index 26fcb39..552078e 100644 --- a/easy_tpp/model/ithp.py +++ b/easy_tpp/model/ithp.py @@ -128,8 +128,7 @@ class ITHP(BaseModel): generic ``num_heads`` setting (whose default is 2). * Paper Eq. (5) scales the 2M-wide query/key dot product by sqrt(2M). Public ``dynamic_v_attention`` instead divides by sqrt(d_k) - (``SubLayers.py`` line 222). We retain the executed public rule so the - released architecture and our saved checkpoints have the same scores. + (``SubLayers.py`` line 222). We follow the executed public rule. * Paper Eq. (7) writes a type-specific softplus head directly on the attention-weighted value sum. However, paper Sec. 5.4 explicitly says the encoder keeps a skip connection and requires M_V=2M; Eq. (7) does @@ -157,9 +156,10 @@ class ITHP(BaseModel): * For that integral, ``integration_method='mc'`` draws independent uniform times per interval by default; ``'trapezoid'`` uses evenly spaced nodes and ``'fixed_grid'`` uses per-interval midpoint cells. - This replaces the authors' active 0.1 grid, but retains the stochastic - sampler of our original EasyTPP runs. The separate public MC helper - (``Utils.py`` lines 42-59) is not called by the active ``Main.py`` path. + This replaces the authors' active batch-wide 0.1 grid with a + per-interval approximation independent of other batch members. The + separate public MC helper (``Utils.py`` lines 42-59) is not called by + the active ``Main.py`` path. Our reported runs used 20 draws per interval for training/validation and 100 for test likelihood; sample counts are an approximation choice, not a number specified by paper Sec. 4.5. Even ``fixed_grid`` with @@ -182,8 +182,7 @@ class ITHP(BaseModel): Our reported runs used Adam epsilon 1e-5 and gradient-norm clipping at 1. EasyTPP's standard runner does not apply those settings, so fresh training - here does not exactly reproduce that training path. Saved checkpoints load - unchanged. + here does not exactly reproduce that training path. Pinned code for the line references above: https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process/blob/5db1bb78f3323667e2cef478e177cd35971c4b43/transformer/Models.py @@ -191,7 +190,6 @@ class ITHP(BaseModel): https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process/blob/5db1bb78f3323667e2cef478e177cd35971c4b43/transformer/SubLayers.py https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process/blob/5db1bb78f3323667e2cef478e177cd35971c4b43/Main.py https://github.com/waystogetthere/Interpretable-Transformer-Hawkes-Process/blob/5db1bb78f3323667e2cef478e177cd35971c4b43/Utils.py - Our prior MC sampler: https://github.com/andrewwarrington/HHP/blob/76fbcf00c33f3d2d8e937c5984e9a1e17b218cde/EasyTPP/easy_tpp/model/torch_model/torch_basemodel.py#L170-L197 """ SUPPORTED_INTEGRATION_METHODS = {"mc", "trapezoid", "fixed_grid"} @@ -292,7 +290,7 @@ def make_dtime_loss_samples(self, time_delta_seq): """Draw random MC times or fixed trapezoid nodes per interval. Upstream EasyTPP's BaseModel currently returns ``linspace`` for both - modes; overriding it preserves the stochastic rule of our runs. + modes; overriding it implements the selected integration method. """ if self.use_mc_samples: ratios = torch.rand(