A
Recurrent State-Space Model trained for tutoring state prediction, part of the
KAT system by
Progga AI .
This is a complete world model for predicting tutoring session dynamics — student state transitions, reward signals, and session termination. It uses a DreamerV3-inspired RSSM architecture with VL-JEPA-style EMA target encoding.
TutoringRSSM (2,802,838 params)
├── ObservationEncoder: obs_dim(20) → encoder_hidden(256) → latent_dim(128)
├── ActionEmbedding: action_dim(8) → embed_dim(32)
├── DeterministicTransition: GRU(hidden_dim=512)
├── StochasticLatent: Diagonal Gaussian prior/posterior (latent_dim=128)
├── ObservationDecoder: feature_dim(640) → decoder_hidden(256) → obs_dim(20)
├── RewardPredictor: feature_dim(640) → 1
├── DonePredictor: feature_dim(640) → 1
└── EMATargetEncoder: momentum=0.996 (VL-JEPA heritage)
Training converged smoothly over 100 epochs with consistent eval loss improvement. No catastrophic forgetting or training instability observed.
1 import torch
2 from architecture import TutoringRSSM , TutoringWorldModelConfig
3
4 # Load model
5 config = TutoringWorldModelConfig (
6 obs_dim = 20 , action_dim = 8 ,
7 latent_dim = 128 , hidden_dim = 512 ,
8 encoder_hidden = 256 , decoder_hidden = 256 ,
9 )
10 model = TutoringRSSM ( config ) . cuda ( )
11
12 ckpt = torch . load ( "tutoring_rssm_best.pt" , map_location = "cuda" )
13 model . load_state_dict ( ckpt [ "model_state_dict" ] )
14 model . eval ( )
15
16 # Initialize state
17 h , z = model . initial_state ( batch_size = 1 )
18
19 # Observe a tutoring step
20 obs = torch . randn ( 1 , 20 ) . cuda ( ) # Student observation
21 action = torch . tensor ( [ 0 ] ) . cuda ( ) # SOCRATIC strategy
22 result = model . observe_step ( h , z , action , obs )
23
24 h_new , z_new = result [ "h" ] , result [ "z" ]
25 pred_obs = result [ "pred_obs" ] # Predicted next observation
26 pred_reward = result [ "pred_reward" ] # Predicted reward
27 pred_done = result [ "pred_done" ] # Predicted session end
28
29 # Imagination (planning without observation)
30 imagined = model . imagine_step ( h_new , z_new , torch . tensor ( [ 3 ] ) . cuda ( ) )
31 # Returns predicted state without requiring real observation