Social-JEPA · Animation study

Learning to imagine without text

How Social-JEPA learns interaction dynamics in latent space, fully self-supervised. Follow the sequence, or select a step to inspect it.

Social-JEPA training objective A frozen RoBERTa encoder and a trainable projection map the dialogue history to a latent state. A predictor rolls it forward five steps, conditioned on actions and time gaps. The same frozen encoder and an exponential-moving-average copy of the projection map the real future histories to target states with no gradient. A loss pulls each predicted state toward its target, with VICReg regularisation; gradients update only the projection and predictor. TARGETS from the conversation’s real future PREDICTIONS rolled forward from the present gradients Ht+k real future RoBERTa-base frozen · shared πψ EMA copy stop-gradient Ht dialogue so far RoBERTa-base frozen · 125M πφ trainable st agent action A and time gap Δt condition every step EMA update ψ ← τψ + (1 − τ)φ VICReg: variance + covariance on s ℒJEPA per-step match, weighted by γk ℒ = Σk γk ‖ŝt+k − sg(stgtt+k)‖2 + λvar ℒvar + λcov ℒcov

A frozen RoBERTa-base encoder (125M parameters) reads the dialogue history Ht, keeping its first and most recent tokens so that stable traits and the current state both fit in context. A trainable projection πφ maps it to a 64-dimensional state st.

The predictor Pθ, a 3-layer MLP, rolls the state forward one turn at a time, conditioned on embeddings of the action the agent actually took, A, and the time gap Δt before the next reply. No text is generated at any step.

Training targets come from the conversation’s real future. The same frozen encoder reads each future history Ht+k, and πψ, an exponential moving average of πφ, maps it to a target state. Targets receive no gradient.

Each prediction is pulled toward its target, discounted by γk over K = 5 steps, while VICReg variance and covariance terms keep the embedding from collapsing. Gradients update only πφ and Pθ (about 500K parameters); πψ then tracks πφ by moving average.

Step 1 of 4
Frozen Trainable / latent state Training signal ds = 64, K = 5, γ = 0.95, λvar = 1.0, λcov = 0.04

Training uses only observed text, actions, and time gaps; LID-Bench’s oracle latent states are never seen.