PennyLaneAI / PennyLaneAI/catalyst

[BUG] Passing observables as parameters triggers an exception

Open
#339 2 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

bug
Dominant language
Python
Stars
234
Forks
84
Avg merge
2d 15h
Merged PRs (30d)
66

Description

import pennylane as qml
import pytest
from catalyst import qjit


def test_observable_as_parameter(backend):
    """Test to see if we can pass an observable parameter to qfunc."""

    coeffs0 = [0.3, -5.1]
    H0 = qml.Hamiltonian(qml.math.array(coeffs0), [qml.PauliZ(0), qml.PauliY(1)])

    @qjit
    def circuit(obs):
        return qml.expval(obs)

    circuit(H0)

Taking H0 as an observable will cause line 254 in compilation_pipelines.py to fail.
I suspect this is related to pytrees.

237     @staticmethod
238     def get_runtime_signature(*args):
239         """Get signature from arguments.
240 
241         Args:
242             *args: arguments to the compiled function
243 
244         Returns:
245             a list of JAX shaped arrays
246         """
247         args_data, args_shape = tree_flatten(args)
248 
249         try:
250             r_sig = []
251             for arg in args_data:
252                 r_sig.append(jax.api_util.shaped_abstractify(arg))
253             # Unflatten JAX abstracted args to preserve the shape
254             return tree_unflatten(args_shape, r_sig)
255         except Exception as exc:
256             arg_type = type(arg)
257             raise TypeError(f"Unsupported argument type: {arg_type}") from exc

This is the exception that is triggered (before being caught immediately after in line 255):

TypeError: float() argument must be a string or a real number, not 'ShapedArray'

This is because the Hamiltonian unflatten function will attempt to build a Hamiltonian object with a ShapedArray.

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start in compilation_pipelines.py at get_runtime_signature, then reproduce the failure with the observable-as-parameter example in the issue. Inspect how tree_flatten and tree_unflatten rebuild the Hamiltonian; done means the qjit circuit accepts H0 as an observable without raising the ShapedArray conversion exception.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
compilers
Issue type
Bug
Difficulty
3/5
Estimated time
1-2 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.