How do I do the same thing for a list of labels?
label = [2,3,5,7,1,0…]
label = 3
NumClass = 10
NumRows = 100
mask = torch.zeros(100, 64)
ones = torch.ones(1, 64)
ElementsPerClass = NumRows//NumClass
mask [ ElementsPerClass*label : ElementsPerClass*(label+1) ] = ones