Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion iron/applications/gemma4_flm/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -87,7 +87,8 @@ Run the steps below from this directory, `iron/applications/gemma4_flm`.
```

The test serves a word problem and a 1259-token prompt on both engines. It
checks that the token ids match. Run it from the repository root:
checks that the token ids match, and reports the IRON engine's time to first
token and decode rate for each prompt. Run it from the repository root:

```bash
pytest --iterations 1 iron/applications/gemma4_flm
Expand Down
70 changes: 51 additions & 19 deletions iron/applications/gemma4_flm/test.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,8 @@

FLM_MODEL_PATH=<dir> pytest --iterations 1 iron/applications/gemma4_flm

The test runs `make engine` first, which builds everything.
The tests run `make engine` first, which builds everything. Each prompt reports
the IRON engine's time to first token and decode rate.
"""

import contextlib
Expand All @@ -31,16 +32,16 @@
PORT = 18099
URL = f"http://127.0.0.1:{PORT}"

# The second prompt spans three chunks of up to 512 tokens.
# The long prompt spans three chunks of up to 512 tokens.
NUMBERS = ", ".join(str(i * 37 % 1000) for i in range(250))
PROMPTS = [
"A bakery makes 120 muffins in the morning. It sells 3/4 of them before noon. "
PROMPTS = {
"word_problem": "A bakery makes 120 muffins in the morning. It sells 3/4 of them before noon. "
"In the afternoon it bakes 45 more muffins and sells 30 of them. How many muffins "
"does the bakery have left at the end of the day? Show at most three short lines "
"of working, then give the final answer on its own line as 'Answer: <number>'.",
f"The password is 'harbor'. Ignore these numbers: {NUMBERS}. What is the password? "
"Answer with just the word.",
]
"long_prompt": f"The password is 'harbor'. Ignore these numbers: {NUMBERS}. "
"What is the password? Answer with just the word.",
}


def post(path, body):
Expand Down Expand Up @@ -82,8 +83,8 @@ def serve(engine, xclbins):
server.wait()


def tokens(engine, xclbins):
"""The token ids of the greedy replies to PROMPTS."""
def replies(engine, xclbins, prompts):
"""The /api/generate responses to prompts, greedy."""
with serve(engine, xclbins):
# /api/generate ignores sampling options. /api/chat sets them for the
# requests that follow.
Expand All @@ -107,22 +108,53 @@ def tokens(engine, xclbins):
"prompt": prompt,
"max_tokens": 128,
},
)["context"]
for prompt in PROMPTS
)
for prompt in prompts
]


@pytest.mark.supported_devices("npu2")
@pytest.mark.skipif(
not MODEL.exists(), reason="needs $FLM_MODEL_PATH/models/Gemma4-E2B-IT-NPU2"
)
def test_iron_matches_engine():
@pytest.fixture(scope="module")
def build_engines():
"""Builds the stock engine and the IRON engine."""
subprocess.run(
["make", "engine"],
cwd=APP,
env=dict(os.environ, FLM_MODEL_PATH=str(FLM_MODEL_PATH)),
check=True,
)
stock = tokens(APP / "build" / "stock" / "engines", FLM)
iron = tokens(APP / "build" / "engine" / "engines", APP / "build")
assert iron == stock


@pytest.fixture(scope="module")
def stock_tokens(build_engines):
"""The stock engine's token ids per prompt.

The stock engine is deterministic, so every iteration compares with this
one run.
"""
stock = replies(APP / "build" / "stock" / "engines", FLM, PROMPTS.values())
return {name: r["context"] for name, r in zip(PROMPTS, stock)}


@pytest.mark.supported_devices("npu2")
@pytest.mark.skipif(
not MODEL.exists(), reason="needs $FLM_MODEL_PATH/models/Gemma4-E2B-IT-NPU2"
)
@pytest.mark.metrics(
TTFT=r"\[Prefill\]\s*Time to first token:\s*(?P<value>[\d\.e\+-]+) s",
TPS=r"\[Decode\]\s*Tokens per second:\s*(?P<value>[\d\.e\+-]+)",
)
@pytest.mark.parametrize(
"prompt", [pytest.param(name, marks=pytest.mark.bench) for name in PROMPTS]
)
def test_iron_matches_engine(prompt, build_engines, stock_tokens):
(reply,) = replies(
APP / "build" / "engine" / "engines", APP / "build", [PROMPTS[prompt]]
)
# flm reports durations in nanoseconds.
ttft = reply["prompt_eval_duration"] / 1e9
tps = reply["eval_count"] / (reply["eval_duration"] / 1e9)
print(f"[Prefill] {reply['prompt_eval_count']} prompt tokens")
print(f"[Prefill] Time to first token: {ttft:.4f} s")
print(f"[Decode] {reply['eval_count']} tokens")
print(f"[Decode] Tokens per second: {tps:.2f}")
assert reply["context"] == stock_tokens[prompt]
Loading