genie_coinrun_world_model(pretrained: bool = False, overrides: object = {})Construct the CoinRun case study as a playable environment.
Model Size
Parameters
pretrainedbool= FalseNo weights were released;
True raises.**overridesobject= {}Optional
GenieConfig field overrides.Returns
GenieForWorldModelingThe networks, plus MaskGIT decoding and the rollout that plays them.
Notes
Reference: Bruce et al., arXiv:2402.15391, Appendix F, Table 17: 25 MaskGIT steps at temperature 1.
Examples
>>> import lucid
>>> from lucid.models import genie_coinrun_world_model
>>> model = genie_coinrun_world_model(
... sample_size=16, num_frames=3, num_codes=16, code_dim=4,
... tokenizer_encoder_layers=1, tokenizer_encoder_dim=16,
... tokenizer_encoder_heads=2, tokenizer_decoder_layers=1,
... tokenizer_decoder_dim=16, tokenizer_decoder_heads=2,
... action_patch_size=8, action_dim=4,
... action_encoder_layers=1, action_encoder_dim=16, action_encoder_heads=2,
... action_decoder_layers=1, action_decoder_dim=16, action_decoder_heads=2,
... dynamics_layers=1, dynamics_dim=16, dynamics_heads=2,
... maskgit_steps=3).eval()
>>> model.config.num_latent_actions
6
>>> out = model(lucid.rand(2, 2, 3, 16, 16), lucid.tensor([[5], [0]], dtype=lucid.int64))
>>> out.frames.shape, out.tokens.shape
((2, 1, 3, 16, 16), (2, 3, 16))