LLM Inference Simulators — Presentation 09

From PyTorch, ONNX & HEIR to a Simulator

How real applications reach a simulated accelerator: graph capture with torch.fx / torch.export and torch.compile backends, out-of-tree PyTorch devices, ONNX Runtime execution providers, MLIR, and Google's HEIR compiler for fully homomorphic encryption.

torch.export torch.compile PrivateUse1 ONNX Runtime EP MLIR HEIR / FHE
Model → Graph → Lowering → Op trace → Simulator
00

Topics We'll Cover

Concepts used here, and where they are explained. Each links to its series glossary entry: a short explanation, then links to the slides that explain it in depth, in this series, the Local LLM Hosting and Key Publications decks, or the FHE Accelerator Simulators series.

01

Why Integration Is Half the Job

A simulator that only runs hand-written workloads answers the questions its authors thought of. One that runs what users actually write (PyTorch models, ONNX graphs, FHE programs) answers the questions customers will ask. "Real-world applications executed on the simulator" means the simulator sits at the bottom of a compiler toolchain, exactly where the silicon will.

Frontends
PyTorch · ONNX · HEIR
→
Graph IR
FX / ATen · ONNX · MLIR
→
Lowering & mapping
ops → hardware primitives
→
Simulator
timing ± function
→
Metrics back
latency, hot-spots, traces
02

Three Ways to Drive a Simulator

ApproachHowStrengthsWeaknesses
Ahead-of-time graphExport the model to a graph (torch.export, ONNX, MLIR), lower it, simulate the operator scheduleWhole-program view; compiler-style optimisation and fusion; no runtime neededDynamic control flow and data-dependent shapes are awkward
Runtime traceRun the model (even on fake tensors), intercept every operator, record name and shapes, replay into the simulatorWorks on almost any model; captures what really executesOne trace per input shape; no compiler-level transformations
Simulator as a deviceRegister the simulator as a framework backend; ops dispatched to it compute results and accumulate simulated timeUsers change one line; functional correctness can be checked; drives software bring-upMost engineering; speed limited by functional computation

Most teams build all three over time: traces first (cheap), graphs next (for the compiler), and a device last (for bring-up and customers).

03

PyTorch I: Graph Capture with torch.export

torch.export traces a model into an ExportedProgram: a single FX graph of ATen-level operators with shape metadata on every node. It is the modern, sound replacement for TorchScript as an export format.

Turn a PyTorch model into an operator list for the cost model
import torch
from torch.export import export

ep = export(model, (torch.randn(1, 2048, 4096),))
ep = ep.run_decompositions()       # lower to Core ATen: linear becomes permute + mm, etc.
ops = []
for node in ep.graph.nodes:
    if node.op == "call_function":
        ins  = [a.meta["val"].shape for a in node.args if isinstance(a, torch.fx.Node)]
        out  = node.meta["val"]
        ops.append((node.target, ins, getattr(out, "shape", None)))
# e.g. (aten.mm.default, [(2048, 4096), (4096, 14336)], (2048, 14336))
# without run_decompositions() you would see aten.linear with (14336, 4096) weights
schedule = mapper.lower(ops)        # your hardware mapping
report   = simulator.run(schedule)  # your simulator
04

PyTorch II: A torch.compile Backend

torch.compile (TorchDynamo) captures graphs from ordinary Python as it runs and hands each graph to a backend: any callable that takes an FX GraphModule and example inputs and returns a callable. That is a clean hook for a simulator.

A simulator backend: estimate on compile, then run eagerly for correct outputs
def sim_backend(gm: torch.fx.GraphModule, example_inputs):
    ops = [(n.target, n.meta["example_value"].shape)      # Dynamo's fake-tensor metadata
           for n in gm.graph.nodes if n.op == "call_function"]
    est = simulator.estimate(ops)                    # latency, utilisation, hot-spots
    log.info("graph of %d ops: %.2f ms simulated", len(ops), est.ms)
    return gm.forward                               # functional result from eager PyTorch

model = torch.compile(model, backend=sim_backend)
model(x)                                             # user code unchanged
05

PyTorch III: Dispatch Interception and the Meta Device

Every PyTorch operator passes through the dispatcher. A TorchDispatchMode sees each ATen call with its arguments, which is enough to record an exact operator trace of any model. Combine it with the meta device (tensors with shapes but no storage) and you can trace a 70B model on a laptop, with no weights and no arithmetic.

Trace every operator of a model that would never fit in memory
from torch.utils._python_dispatch import TorchDispatchMode

class OpTrace(TorchDispatchMode):
    def __init__(self):
        super().__init__(); self.ops = []
    def __torch_dispatch__(self, func, types, args=(), kwargs=None):
        out = func(*args, **(kwargs or {}))
        shapes = [tuple(a.shape) for a in args if isinstance(a, torch.Tensor)]
        self.ops.append((str(func), shapes))
        return out

with torch.device("meta"):
    model = LlamaForCausalLM(config)               # 70B parameters, zero bytes allocated
    x = torch.zeros(1, 2048, dtype=torch.long)
with OpTrace() as t:
    model(x)
simulator.replay(t.ops)                            # a real op sequence, simulated

Checked with PyTorch 2.14 (CPU): a 70B-sized MLP block traced on the meta device yields aten.mm on (2048, 8192) × (8192, 28672) with nothing allocated. This is the cheapest route to "real models on the simulator", and the trace doubles as a regression input: if a library upgrade changes the operator sequence, the diff shows exactly how.

06

PyTorch IV: The Simulator as a Device

PyTorch reserves a dispatch key, PrivateUse1, for out-of-tree accelerators. A vendor renames it (for example torch.utils.rename_privateuse1_backend("photon")), registers kernels for it in C++, and users write model.to("photon"). Several production accelerators integrate this way.

C++: register a kernel for the out-of-tree device
at::Tensor photon_mm(const at::Tensor& a, const at::Tensor& b) {
    sim::Stats::instance().add("mm", a.sizes(), b.sizes());    // timing model
    return photon::functional_mm(a, b);                     // bit-accurate result (incl. precision effects)
}
TORCH_LIBRARY_IMPL(aten, PrivateUse1, m) {
    m.impl("mm", photon_mm);
}

Why a functional device matters for novel compute

Optical and analogue compute have finite precision and noise. A device-level simulator can apply the hardware's numerical behaviour to real tensors, so users see whether a model's accuracy survives, not only how fast it runs.

The cost

Memory allocators, copies, streams and a long tail of operators. Start with a fallback to CPU for unsupported operators and grow coverage, tracking coverage as an engineering metric.

07

ONNX and ONNX Runtime

ONNX is a framework-neutral graph format: nodes with standard operator types, versioned opsets, typed tensors. ONNX Runtime (ORT) executes ONNX graphs through pluggable Execution Providers (EPs) such as CPU, CUDA, TensorRT, OpenVINO and QNN.

How an EP plugs in

  • ORT asks each EP, in priority order, which nodes it can run (GetCapability).
  • The graph is partitioned; each EP's subgraphs may be fused and compiled into single kernels.
  • Anything unclaimed falls back to the CPU EP.
  • A simulator EP claims the nodes the future hardware supports, and the partition itself becomes a metric: how much of each model runs on the accelerator.

ORT already emits Chrome traces

so = ort.SessionOptions()
so.enable_profiling = True
sess = ort.InferenceSession("model.onnx", so,
         providers=["CPUExecutionProvider"])
sess.run(None, {"input": x})
path = sess.end_profiling()  # Chrome-trace JSON

Same format as the simulator's traces (deck 06): baseline and simulated timelines open side by side in Perfetto.

Walk an ONNX graph with inferred shapes
import onnx
from onnx import shape_inference
m = shape_inference.infer_shapes(onnx.load("model.onnx"))
shapes = {v.name: [d.dim_value for d in v.type.tensor_type.shape.dim]
          for v in [*m.graph.input, *m.graph.value_info, *m.graph.output]}
shapes.update({t.name: list(t.dims) for t in m.graph.initializer})   # weights
for node in m.graph.node:
    attrs = {a.name: onnx.helper.get_attribute_value(a) for a in node.attribute}
    simulator.add_op(node.op_type, [shapes.get(i) for i in node.input], attrs)
08

MLIR in One Slide

MLIR (Multi-Level Intermediate Representation, part of LLVM) is a framework for building compilers out of dialects, families of operations at a given abstraction, connected by progressive lowering passes.

High level
torch / StableHLO / TOSAsecret (HEIR)
Domain
linalg / tensorckks / bgv / cggi (HEIR)
Structural
scf / affine / memrefpolynomial / mod_arith (HEIR)
Target
your accelerator's dialectllvm
09

Fully Homomorphic Encryption in One Slide

FHE lets a server compute on encrypted data without decrypting it. The client encrypts, the server evaluates a function on ciphertexts, and only the client can decrypt the result. It is the strongest form of confidential computing, and very expensive.

The arithmetic

  • Modern schemes (BGV, BFV, CKKS for approximate arithmetic, CGGI/TFHE for Boolean circuits) are built on the Ring-LWE problem.
  • A ciphertext is a pair of polynomials in Zq[X]/(XN+1), with N typically 213–217, and q split into many machine-word primes (RNS limbs).
  • Polynomial multiplication uses the NTT, the finite-field cousin of the FFT, in O(N log N).
  • Multiplications and rotations need key switching, which reads large evaluation keys.
  • Noise grows with each operation; bootstrapping refreshes it, at a cost of thousands of NTTs.

Why it needs hardware

  • FHE is commonly cited as several orders of magnitude slower than computing on plaintext on CPUs.
  • The work is dominated by NTTs, element-wise modular multiply-adds, and data movement of ciphertexts and keys.
  • Accelerator research (GPU, FPGA, ASIC, and photonic proposals) targets exactly those primitives.
  • Transforms are a natural fit for optics: a lens performs a Fourier transform physically. The challenge is mapping exact modular NTT arithmetic onto analogue optical transforms with enough precision, which is precisely the kind of question a full-system simulator must answer.
10

HEIR: An MLIR Compiler for FHE

HEIR (Homomorphic Encryption Intermediate Representation) is an open-source, MLIR-based FHE compiler toolchain led by Google. Its aim is to let developers write ordinary-looking programs, mark what is secret, and compile them to an FHE scheme and a backend: a software library or a hardware accelerator.

Program
secret-annotated MLIR, or a Python frontend
→
Scheme dialects
ckks / bgv / cggi, lwe
→
Arithmetic
polynomial, mod_arith, RNS
→
Backends
OpenFHE, Lattigo, …, accelerators
11

Interactive: FHE Data-Movement Calculator

A first-order look at why FHE accelerators are memory and data-movement machines. A ciphertext is 2 polynomials × N coefficients × L limbs × 8 bytes. This is a rough sizing aid, not a scheme-accurate cost model.

65536
30
200
1000
Ciphertext size
—
NTT butterflies / transform
—
Time to transform 1 ciphertext
—
Time to read 1 ciphertext
—

For the full treatment (bootstrapping, key traffic, scratchpad sizing and optical NTT engines, simulated), see the sister series FHE Accelerator Simulators. The same lesson as the LLM roofline in deck 03: a faster compute primitive moves the bottleneck to memory and data movement unless the architecture keeps data close to the engine. Key-switching keys (often far larger than a ciphertext) make this worse.

12

Putting the Toolchain Together

LayerAI pathFHE pathOwned by
FrontendPyTorch (export, compile, dispatch), ONNXHEIR frontendsFramework integration
IR and loweringFX/ATen → MLIR (torch-mlir, StableHLO, linalg)secret → ckks/bgv → polynomialCompiler
MappingTile GEMMs, schedule attention, place KVSchedule NTTs, key switches, rotationsCompiler and architecture
SimulatorTiming (and optionally function) of the scheduleSame engine, FHE primitivesSimulation
OraclePyTorch eager outputsOpenFHE or Lattigo outputsVerification
ReportLatency, throughput, hot-spots, tracesLatency per operation, bootstrapping cost, data movementSimulation
One engine, many workloads

The discrete-event core, metrics, traces and test ladder from decks 02, 06 and 08 do not care whether the operations are GEMMs or NTTs. Keep the cost models and workload frontends pluggable and one simulator serves AI and FHE markets alike.

13

What to Take Away

Next

Deck 10 is a reading list for going deeper on everything in this series.