mini-deepseek-v4-flash / validation /tests_fuse2_cache_modal.py
Akahsizrr's picture
Publish compatibility source validation/tests_fuse2_cache_modal.py
db87ef5 verified
Raw
History Blame Contribute Delete
2.11 kB
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))