from __future__ import annotations import json import modal image = ( modal.Image.debian_slim(python_version="3.11") .pip_install("torch==2.7.0", "transformers==5.14.1") .add_local_file("microscope/fuse2_model.py", "/root/fuse2_model.py") ) app = modal.App("fuse2-cache-test") @app.function(image=image, cpu=4, memory=8192, timeout=600) def run(): import sys import torch from transformers import Qwen3Config sys.path.insert(0, "/root") from fuse2_model import Fuse2Config, Fuse2ForCausalLM, Fuse2AugmentedLayer torch.manual_seed(7) config = Fuse2Config( vocab_size=128, hidden_size=64, intermediate_size=128, num_hidden_layers=2, num_attention_heads=4, num_key_value_heads=2, head_dim=16, max_position_embeddings=64, experts_per_layer={"0": [0], "1": [0]}, expert_hidden_size=64, expert_intermediate_size=32, top_k_experts=1, pad_token_id=0, bos_token_id=1, eos_token_id=2, ) model = Fuse2ForCausalLM(config).eval() for layer in model.model.layers: if isinstance(layer, Fuse2AugmentedLayer): torch.nn.init.normal_(layer.bridge_out.weight, std=0.02) torch.nn.init.normal_(layer.repair_up.weight, std=0.02) ids = torch.tensor([[5, 9, 13, 17, 21, 25]], dtype=torch.long) with torch.no_grad(): full = model(input_ids=ids, use_cache=False, return_dict=True).logits[:, -1] prefix = model(input_ids=ids[:, :-1], use_cache=True, return_dict=True) cached = model( input_ids=ids[:, -1:], past_key_values=prefix.past_key_values, use_cache=True, return_dict=True, ).logits[:, -1] max_error = (full.float() - cached.float()).abs().max().item() result = {"max_logit_error": max_error, "cache_type": type(prefix.past_key_values).__name__} if max_error > 2e-3: raise AssertionError(json.dumps(result)) return result @app.local_entrypoint() def main(): print(json.dumps(run.remote(), indent=2))