Overloading __iter__ for a Python dataclass with fields containing list of another dataclass

Viewed 583

I don't know how to properly overload the iterator for my WholeImage class.

from dataclasses import dataclass, astuple
from typing import List
import numpy as np

@dataclass(frozen=True)
class ActImage:
    __slots__ = ['name', 'action']
    name: np.ndarray
    action: np.ndarray

    def __iter__(self):
        return iter(astuple(self))


@dataclass(frozen=True)
class WholeImage:
    __slots__ = ['value', 'x_axis', 'y_axis', 'act']
    value: np.ndarray
    x_axis: List[np.ndarray]
    y_axis: List[np.ndarray]
    act: List[ActImage]

    def __iter__(self):
        yield from (self.value, *self.x_axis, *self.y_axis, *self.act)

The way it works now, it returns numpy ndarrays for value and elements in x_axis and y_axis lists, but ActImage instances for elements in act list.

The way I'd like it to behave, is to return all ndarrays, including the "nested" ones in ActImage.name and ActImage.action, as in the pseudocode below:

def __iter__(self):
   yield from (self.value, *self.x_axis, *self.y_axis, 
               self.act[x].name, self.act[x].action # for x in range(len(act))
               ) 

I'd like to have my code as generic as possible, because new fields (of np.ndarray type too) may be added to ActImage class in the future.

0 Answers
Related