# Codeblock 8 def class_loss_test(): target = torch.zeros(BATCH_SIZE, S, S, (C+5)) prediction = torch.zeros(BATCH_SIZE, S, S, (C+B*5)) target[0, 3, 3, 21:25] = torch.tensor([0.4, 0.5, 2.4, 3.2]) target[0, 3, 3, 20] = 1.0 target[0, 3, 3, 7] = 1.0 prediction[0, 3, 3, 21:25] = torch.tensor([0.4, 0.5, 2.4, 3.2]) prediction[0, 3, 3, 7] = 1.0 #(1) #prediction[0, 3, 3, 7:9] = torch.tensor([0.9, 0.1]) #(2) #prediction[0, 3, 3, 7:9] = torch.tensor([0.2, 0.8]) #(3) target = target.reshape(BATCH_SIZE, S*S*(C+5)) prediction = prediction.reshape(BATCH_SIZE, S*S*(C+B*5)) class_loss = loss(target, prediction)[3] return class_loss class_loss_test()