Skip to content

simulate

Circuit simulation methods.

simulate(circuit, initial_state, mode=SimulateMode.DEFAULT, save_states=True, **kwargs)

Simulate a circuit and optionally retain each layer's states.

Parameters:

Name Type Description Default
circuit Circuit

Circuit to simulate.

required
initial_state Qarray

Initial ket or density matrix.

required
mode SimulateMode

Simulation mode, or each layer's default mode.

DEFAULT
save_states bool

Whether to retain intermediate layer states.

True

Returns:

Type Description
Results

Saved states, or only the final state when save_states=False.

Source code in jaxquantum/circuits/simulate.py
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
def simulate(
    circuit: Circuit,
    initial_state: Qarray,
    mode: SimulateMode = SimulateMode.DEFAULT,
    save_states: bool = True,
    **kwargs,
) -> Results:
    """Simulate a circuit and optionally retain each layer's states.

    Args:
        circuit: Circuit to simulate.
        initial_state: Initial ket or density matrix.
        mode: Simulation mode, or each layer's default mode.
        save_states: Whether to retain intermediate layer states.

    Returns:
        Saved states, or only the final state when ``save_states=False``.
    """

    results = Results.create(
        [_single_state_batch(initial_state)] if save_states else []
    )
    state = initial_state
    start_time = 0
    if not save_states:
        kwargs.setdefault("saveat_tlist", jnp.array([]))

    for layer in circuit.layers:
        result_dict = _simulate_layer(
            layer,
            state,
            mode=mode,
            start_time=start_time,
            **kwargs,
        )
        result = result_dict["result"]
        start_time = result_dict["start_time"]
        state = result[-1]
        if save_states:
            results.append(result)

    if not save_states:
        results.append(_single_state_batch(state))
    return results

simulate_expectations(circuit, initial_state, observables, mode=SimulateMode.DEFAULT, include_initial=True, **kwargs)

Return the final state and per-layer expectation values.

Source code in jaxquantum/circuits/simulate.py
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
def simulate_expectations(
    circuit: Circuit,
    initial_state: Qarray,
    observables: list[Qarray],
    mode: SimulateMode = SimulateMode.DEFAULT,
    include_initial: bool = True,
    **kwargs,
):
    """Return the final state and per-layer expectation values."""
    if not observables:
        raise ValueError("observables must not be empty")

    kwargs.setdefault("saveat_tlist", jnp.array([]))
    state = initial_state
    start_time = 0.0
    values = [_expectations(state, observables)] if include_initial else []
    for layer in circuit.layers:
        output = _simulate_layer(layer, state, mode, start_time, **kwargs)
        state = output["result"][-1]
        start_time = output["start_time"]
        values.append(_expectations(state, observables))
    if not values:
        return state, _expectations(state, observables)[None][:0]
    return state, jnp.stack(values)

simulate_final(circuit, initial_state, mode=SimulateMode.DEFAULT, **kwargs)

Return only the final circuit state.

Source code in jaxquantum/circuits/simulate.py
335
336
337
338
339
340
341
342
343
def simulate_final(
    circuit: Circuit,
    initial_state: Qarray,
    mode: SimulateMode = SimulateMode.DEFAULT,
    **kwargs,
) -> Qarray:
    """Return only the final circuit state."""
    kwargs.setdefault("saveat_tlist", jnp.array([]))
    return _evolve_circuit(circuit, initial_state, mode, **kwargs)[0]

simulate_repeated(circuit, initial_state, repetitions, mode=SimulateMode.DEFAULT, **kwargs)

Apply one circuit repeatedly with a compiled loop.

Source code in jaxquantum/circuits/simulate.py
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
def simulate_repeated(
    circuit: Circuit,
    initial_state: Qarray,
    repetitions: int,
    mode: SimulateMode = SimulateMode.DEFAULT,
    **kwargs,
) -> Qarray:
    """Apply one circuit repeatedly with a compiled loop."""
    if repetitions < 0:
        raise ValueError("repetitions must be non-negative")
    if repetitions == 0:
        return initial_state

    kwargs.setdefault("saveat_tlist", jnp.array([]))
    state, start_time = _evolve_circuit(
        circuit,
        initial_state,
        mode,
        **kwargs,
    )

    def repeat(_, carry):
        return _evolve_circuit(circuit, carry[0], mode, carry[1], **kwargs)

    return lax.fori_loop(
        1,
        repetitions,
        repeat,
        (state, start_time),
    )[0]

simulate_repeated_expectations(circuit, initial_state, repetitions, observables, mode=SimulateMode.DEFAULT, include_initial=True, **kwargs)

Return the final state and per-repetition expectation values.

Source code in jaxquantum/circuits/simulate.py
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
def simulate_repeated_expectations(
    circuit: Circuit,
    initial_state: Qarray,
    repetitions: int,
    observables: list[Qarray],
    mode: SimulateMode = SimulateMode.DEFAULT,
    include_initial: bool = True,
    **kwargs,
):
    """Return the final state and per-repetition expectation values."""
    if repetitions < 0:
        raise ValueError("repetitions must be non-negative")
    if not observables:
        raise ValueError("observables must not be empty")
    if repetitions == 0:
        values = _expectations(initial_state, observables)[None]
        return initial_state, values if include_initial else values[:0]

    kwargs.setdefault("saveat_tlist", jnp.array([]))
    state, start_time = _evolve_circuit(
        circuit,
        initial_state,
        mode,
        **kwargs,
    )
    first = _expectations(state, observables)

    def repeat(carry, _):
        state, start_time = _evolve_circuit(
            circuit,
            carry[0],
            mode,
            carry[1],
            **kwargs,
        )
        return (state, start_time), _expectations(state, observables)

    (state, _), rest = lax.scan(
        repeat,
        (state, start_time),
        None,
        length=repetitions - 1,
    )
    values = jnp.concatenate((first[None], rest), axis=0)
    if include_initial:
        values = jnp.concatenate(
            (_expectations(initial_state, observables)[None], values),
            axis=0,
        )
    return state, values