How the Phoenix Transformer Model Predicts Engagement Probabilities for User Actions
The Phoenix transformer model predicts engagement probabilities by processing user history and candidate posts through a Transformer-based architecture to generate logits, applying sigmoid activation to obtain action-level probabilities, and aggregating these via maximum pooling across related low-level client actions to produce high-level engagement scores.
The xai-org/x-algorithm repository contains the open-source recommendation engine powering X's content distribution, centered on the Phoenix transformer model. This system estimates the likelihood of diverse user interactions—from clicks and dwells to replies and retweets—by bridging granular client-side events with abstract engagement categories used for ranking and recommendation.
The Transformer Architecture for Logits Generation
The prediction pipeline begins with the RecsysAggregatedModel class, which encapsulates the core transformer logic defined in phoenix/xrex/models/transformer.py. This model receives embeddings representing the user's history, candidate posts, and auxiliary contextual features, then processes them through a stack of multi-head attention layers.
The final hidden state undergoes a linear projection to produce a dense logit tensor (logits). This tensor's last dimension indexes every possible low-level client action type (e.g., specific link clicks, scroll events, or media views). The shape of this output is (batch, seq, num_action_types), providing a raw score for each action across the sequence length.
Converting Logits to Action Probabilities
Once the transformer produces logits, the system converts these to probabilities using the sigmoid function. In phoenix/xrex/models/recsys_model.py (lines 422-426), the implementation casts logits to float32 for numerical stability before applying the activation:
probs = jax.nn.sigmoid(logits) # shape: (batch, seq, num_action_types)
This operation yields a probability between 0 and 1 for each individual low-level action, independent of other actions. The resulting probs tensor serves as the foundation for all downstream engagement predictions.
Aggregating Low-Level Actions into Engagement Categories
The Phoenix model does not predict high-level engagements (like "click" or "view") directly. Instead, it maps abstract engagement categories to collections of specific client actions. This mapping is defined statically in phoenix/xrex/data/recsys/constants.py (lines 59-64) through the engagement_to_ids function:
eng_to_ids = engagement_to_ids(metric_group)
# e.g. {"click": [ACTION_OPEN_LINK, ACTION_LINK_CLICK], ...}
The Maximum Probability Aggregation Strategy
The critical aggregation logic resides in get_probs_and_labels within phoenix/xrex/models/recsys_model.py (lines 22-34). This utility iterates over each engagement category, extracts the probability slice corresponding to its constituent actions, and computes the maximum probability across that group:
for eng_name, client_event_list in eng_to_ids.items():
id_array = jnp.array(client_event_list)
p = probs[..., id_array].max(axis=-1) # max over the actions
p = jnp.clip(p, _CLAMP_eps, 1.0 - _CLAMP_eps) # numerical safety
prob_labels[eng_name] = (p, y) # y = ground‑truth label
This maximum aggregation strategy ensures that if any single low-level action strongly indicates user engagement, the model captures that peak signal. The function also applies clamping using _CLAMP_eps to prevent numerical instability during logarithmic operations in loss computation.
Integration with Training Objectives and Metrics
The resulting dictionary maps each high-level engagement (e.g., "click", "reply", "retweet") to a tuple containing the aggregated probability and the ground-truth label. These per-engagement probabilities feed directly into loss functions such as binary cross-entropy and RCE (Relative Cross-Entropy), as well as evaluation metrics including AUC and NDCG defined in phoenix/xrex/eval/metrics_recsys.py.
During training, the model learns to optimize these aggregated probabilities, effectively training the underlying transformer to recognize patterns that predict the likelihood of each high-level user action.
Summary
- The Phoenix transformer model uses a
RecsysAggregatedModelarchitecture to generate dense logit vectors for every possible client-side action type. - Sigmoid activation (
jax.nn.sigmoid) converts logits to individual action probabilities with shape(batch, seq, num_action_types). - The
engagement_to_idsmapping inconstants.pygroups low-level action IDs into high-level engagement categories. - The
get_probs_and_labelsfunction aggregates probabilities via maximum pooling across related actions and applies numerical clamping for stability. - These per-engagement probabilities drive binary cross-entropy loss and ranking metrics throughout the training and evaluation pipeline.
Frequently Asked Questions
What is the shape of the output tensor from the Phoenix transformer's forward pass?
The output logit tensor has shape (batch, seq, num_action_types), where the final dimension indexes all possible low-level client actions defined in the action space. After sigmoid conversion, the probability tensor maintains this same shape.
How does the model handle multiple low-level actions that map to the same engagement?
The get_probs_and_labels function extracts the probability slice for each engagement category using probs[..., id_array].max(axis=-1), taking the maximum value across all related low-level actions. This ensures the model captures the strongest engagement signal from any constituent action.
Why does Phoenix use maximum aggregation instead of summing or averaging probabilities?
Maximum aggregation prevents signal dilution. If a specific low-level action (like a particular link click type) strongly indicates engagement, averaging would weaken that signal with less relevant actions, while summing could exceed probability bounds or over-weight redundant actions. The maximum operation ensures the peak probability drives the engagement score.
Where are the engagement-to-action mappings defined in the codebase?
The static mappings reside in phoenix/xrex/data/recsys/constants.py (lines 59-64) within the engagement_to_ids function. These mappings organize raw client events into metric groups such as "default" or "video", allowing the system to aggregate diverse signals into standardized engagement categories like "click" or "dwell".
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →