diff --git a/src/coffea/nanoevents/methods/fcc.py b/src/coffea/nanoevents/methods/fcc.py index 6c60b1e77..e714b2f7c 100644 --- a/src/coffea/nanoevents/methods/fcc.py +++ b/src/coffea/nanoevents/methods/fcc.py @@ -418,8 +418,8 @@ def namefcn(self): behavior_edm4hep1[classname].__repr__ = namefcn -@awkward.mixin_class(behavior_edm4hep1) -class MCParticle(edm4hep.MCParticle): # noqa: F811 +@awkward.mixin_class(behavior_edm4hep1, name="MCParticle") +class MCParticle_edm4hep1(edm4hep.MCParticle): """EDM4HEP Datatype: MCParticle; Modified for FCC""" # Get MC Daughters @@ -447,18 +447,18 @@ def get_parents(self, dask_array): ).items(): behavior_edm4hep1.setdefault(_key, _value) del _key, _value -MCParticleArray.ProjectionClass2D = vector.TwoVectorArray # noqa: F821 -MCParticleArray.ProjectionClass3D = vector.ThreeVectorArray # noqa: F821 -MCParticleArray.ProjectionClass4D = MCParticleArray # noqa: F821 -MCParticleArray.MomentumClass = vector.LorentzVectorArray # noqa: F821 -MCParticleRecord.ProjectionClass2D = vector.TwoVectorRecord # noqa: F821 -MCParticleRecord.ProjectionClass3D = vector.ThreeVectorRecord # noqa: F821 -MCParticleRecord.ProjectionClass4D = vector.LorentzVectorRecord # noqa: F821 -MCParticleRecord.MomentumClass = vector.LorentzVectorRecord # noqa: F821 - - -@awkward.mixin_class(behavior_edm4hep1) -class ReconstructedParticle(edm4hep.ReconstructedParticle): # noqa: F811 +MCParticle_edm4hep1Array.ProjectionClass2D = vector.TwoVectorArray # noqa: F821 +MCParticle_edm4hep1Array.ProjectionClass3D = vector.ThreeVectorArray # noqa: F821 +MCParticle_edm4hep1Array.ProjectionClass4D = MCParticle_edm4hep1Array # noqa: F821 +MCParticle_edm4hep1Array.MomentumClass = vector.LorentzVectorArray # noqa: F821 +MCParticle_edm4hep1Record.ProjectionClass2D = vector.TwoVectorRecord # noqa: F821 +MCParticle_edm4hep1Record.ProjectionClass3D = vector.ThreeVectorRecord # noqa: F821 +MCParticle_edm4hep1Record.ProjectionClass4D = vector.LorentzVectorRecord # noqa: F821 +MCParticle_edm4hep1Record.MomentumClass = vector.LorentzVectorRecord # noqa: F821 + + +@awkward.mixin_class(behavior_edm4hep1, name="ReconstructedParticle") +class ReconstructedParticle_edm4hep1(edm4hep.ReconstructedParticle): """EDM4HEP Datatype: Reconstructed particle; Modified for FCC""" # Get MC counterpart @@ -500,18 +500,21 @@ def get_tracks(self, dask_array): _set_repr_name_edm4hep1("ReconstructedParticle") _copy_behaviors("MomentumCandidate", "ReconstructedParticle", behavior_edm4hep1) -ReconstructedParticleArray.ProjectionClass2D = vector.TwoVectorArray # noqa: F821 -ReconstructedParticleArray.ProjectionClass3D = vector.ThreeVectorArray # noqa: F821 -ReconstructedParticleArray.ProjectionClass4D = ReconstructedParticleArray # noqa: F821 -ReconstructedParticleArray.MomentumClass = vector.LorentzVectorArray # noqa: F821 -ReconstructedParticleRecord.ProjectionClass2D = vector.TwoVectorRecord # noqa: F821 -ReconstructedParticleRecord.ProjectionClass3D = vector.ThreeVectorRecord # noqa: F821 -ReconstructedParticleRecord.ProjectionClass4D = vector.LorentzVectorRecord # noqa: F821 -ReconstructedParticleRecord.MomentumClass = vector.LorentzVectorRecord # noqa: F821 - - -@awkward.mixin_class(behavior_edm4hep1) -class ParticleID(edm4hep.ParticleID): # noqa: F811 +_recop_array = ReconstructedParticle_edm4hep1Array # noqa: F821 +_recop_record = ReconstructedParticle_edm4hep1Record # noqa: F821 +_recop_array.ProjectionClass2D = vector.TwoVectorArray +_recop_array.ProjectionClass3D = vector.ThreeVectorArray +_recop_array.ProjectionClass4D = _recop_array +_recop_array.MomentumClass = vector.LorentzVectorArray +_recop_record.ProjectionClass2D = vector.TwoVectorRecord +_recop_record.ProjectionClass3D = vector.ThreeVectorRecord +_recop_record.ProjectionClass4D = vector.LorentzVectorRecord +_recop_record.MomentumClass = vector.LorentzVectorRecord +del _recop_array, _recop_record + + +@awkward.mixin_class(behavior_edm4hep1, name="ParticleID") +class ParticleID_edm4hep1(edm4hep.ParticleID): """EDM4HEP Datatype: ParticleID; Modified for FCC""" # Get ReconstructedParticle @@ -527,8 +530,8 @@ def get_reconstructedparticles(self, dask_array): _set_repr_name_edm4hep1("ParticleID") -@awkward.mixin_class(behavior_edm4hep1) -class Cluster(edm4hep.Cluster): # noqa: F811 +@awkward.mixin_class(behavior_edm4hep1, name="Cluster") +class Cluster_edm4hep1(edm4hep.Cluster): """EDM4HEP Datatype: Cluster; Modified for FCC""" # Get cluster EFlowPhoton @@ -553,8 +556,8 @@ def get_hits(self, dask_array): _set_repr_name_edm4hep1("Cluster") -@awkward.mixin_class(behavior_edm4hep1) -class Track(edm4hep.Track): # noqa: F811 +@awkward.mixin_class(behavior_edm4hep1, name="Track") +class Track_edm4hep1(edm4hep.Track): """EDM4HEP Datatype: Track; Modified for FCC""" # Get Tracks diff --git a/src/coffea/util.py b/src/coffea/util.py index 14672a562..42b880135 100644 --- a/src/coffea/util.py +++ b/src/coffea/util.py @@ -369,6 +369,15 @@ def impl(*args, **kwargs): return descriptor +def _rebuild_dask_property(fget, fset, fdel, doc, dask_get): + prop = _DaskProperty(fget, fset, fdel) + # a doc passed to property() lands in a C slot which, before 3.13, the + # subclass' own docstring shadows -- so assign it to the instance instead + prop.__doc__ = doc + prop._dask_get = dask_get + return prop + + class _DaskProperty(property): _dask_get = None @@ -377,6 +386,16 @@ def dask(self, func): self._dask_get = _make_dask_descriptor(func) return self + def __reduce__(self): + # property keeps fget/fset/fdel in C slots and offers no reduction, so + # a behavior class carrying a dask_property cannot be pickled by value + # (which is what cloudpickle does for classes defined in __main__ or a + # notebook) unless we provide one. + return ( + _rebuild_dask_property, + (self.fget, self.fset, self.fdel, self.__doc__, self._dask_get), + ) + def _adapt_naive_dask_get(func): def wrapper(self, dask_array, *args, **kwargs): diff --git a/tests/test_nanoevents.py b/tests/test_nanoevents.py index cf48c561e..a0ef7edcc 100644 --- a/tests/test_nanoevents.py +++ b/tests/test_nanoevents.py @@ -502,3 +502,36 @@ def test_union_form_genuinely_missing_branch_dask(tmp_path, dask_client): ).events() with pytest.raises(KeyError): events["flag"].compute() + + +def _all_schemas(): + from coffea.nanoevents import schemas + + return [ + getattr(schemas, name) + for name in schemas.__all__ + if hasattr(getattr(schemas, name), "behavior") + ] + + +@pytest.mark.parametrize("schemaclass", _all_schemas(), ids=lambda c: c.__name__) +def test_schema_behavior_survives_pickling(schemaclass): + """Dask ships the schema's behavior dict to distributed workers, so every + entry has to be picklable, and every class in it has to be reachable under + its own qualified name -- otherwise cloudpickle falls back to copying the + class by value, which both bloats the graph and can fail outright. + """ + import sys + + cloudpickle = pytest.importorskip("cloudpickle") + + behavior = dict(schemaclass.behavior()) + cloudpickle.loads(cloudpickle.dumps(behavior)) + + for key, value in behavior.items(): + if isinstance(value, type): + module = sys.modules[value.__module__] + assert getattr(module, value.__qualname__, None) is value, ( + f"{schemaclass.__name__} behavior[{key!r}] is shadowed: " + f"{value.__module__}.{value.__qualname__} is a different class" + ) diff --git a/tests/test_util.py b/tests/test_util.py index f121b0263..01c3959eb 100644 --- a/tests/test_util.py +++ b/tests/test_util.py @@ -71,3 +71,41 @@ def output(x): finally: if os.path.exists(filename): os.remove(filename) + + +def test_dask_property_is_picklable(): + """Behavior classes defined outside an importable module get pickled by + value (e.g. by cloudpickle, when a dask graph goes to a distributed worker), + which walks the class dict -- and plain property objects cannot be pickled. + """ + cloudpickle = pytest.importorskip("cloudpickle") + + from coffea.util import dask_property + + class Thing: + def __init__(self, x): + self.x = x + + @dask_property + def doubled(self): + """twice x""" + return 2 * self.x + + @doubled.dask + def doubled(self, dask_array): + return 20 * dask_array.x + + @dask_property(no_dispatch=True) + def tripled(self): + return 3 * self.x + + # defined in a function body, so this has to go by value + unpickled = cloudpickle.loads(cloudpickle.dumps(Thing)) + + assert unpickled(1).doubled == 2 + assert unpickled(1).tripled == 3 + + doubled = unpickled.__dict__["doubled"] + assert doubled.__doc__ == "twice x" + assert doubled._dask_get(unpickled(1), unpickled, Thing(3)) == 60 + assert unpickled.__dict__["tripled"]._dask_get(unpickled(2), unpickled, None) == 6