Returning indices from pytorch Dataset: Function to alter __getitem__ results in metaclass conflict

Viewed 646

I have multiple classes (for different datasets) that inherit from pytorch's Dataset class. They have a general structure, like so:

from torch.utils.data import Dataset

class SomeDataset(Dataset):

    def __init__(self, data, labels):
        super(SomeDataset, self).__init__()
        self.data = data
        self.labels = labels
        self.__name__ = 'SomeDataset'

    def __getitem__(self, index):
        return {'data': self.data[index], 'label': self.labels[index]}

    def __len__(self):
        return len(data)

Recently I have realised that it would be beneficial to keep track of the labels passed into the Dataloader when batching, so upon googling how to do this I came across this thread, which is where I have adapted the code to write this function:

def return_indices(dataset_class):
    
    def __getitem__(self, index):
        return {'index':1, **dataset_class.__getitem__(self, index)}

    return type(dataset_class.__name__, (dataset_class, ), {'__getitem__': __getitem__})

I had never seen type used like this before, but after some googling, it made some sense, so I tried it out. Unfortunately this led to this error:

TypeError: metaclass conflict: the metaclass of a derived class must be a (non-strict) subclass of the metaclasses of all its bases

which led to a whole lot more googling, and even though I'm beginning to grasp what a metaclass is and how they're used I still can't figure out what is wrong with this approach or how to solve it - and I'm starting to think that maybe it would be easier to rewrite this functionality into my dataset classes instead of having some neat wrapper that does it for me. Can anyone weigh in with whatever it is I'm missing?

1 Answers

Just do this:

def return_indices(dataset_class):
    
    def __getitem__(self, index):
        return {'index':1, **dataset_class.__getitem__(self, index)}
    metacls = type(dataset_class)
    return metacls(dataset_class.__name__, (dataset_class, ), {'__getitem__': __getitem__})

What takes place: as you found out, the 3-paramter call to type is way to create a new-class programatically in Python, without the need for a "class" statement and its body.

But type is the "base metaclass" - and while its instances will be ordinary classes, it also "hardcodes" the metaclass of the class you are creating to itself - in contrast, using the class statement will make Python search for a suitable metaclass among the bases of the class you are creating.

Just using your derived class metaclass (which is obtained by either the one-parameter form of type, as above, or by the __class__ attribute of the class, like in dataset_class.__class__).

Using this as a callable in place of type will have itself as the metaclass, and things should work.

NB: As there are a couple more mechanisms to metaclasses, like __prepare__, just calling the metaclass instead of type will not always work - the correct generic way to do that involves calling types.prepare_class and types.new_class and having a callback to perform the equivalent of the execution of the class body that takes place in the body of a class statement. That won't be needed for most cases.

Related