I am trying to implement the technique called "Invariant risk minimization," which adds a penalty term to the loss function in training machine learning models. The new penalty term's technical definition is the squared gradient norm with respect to a constant classifier. There is an implementation of this "penalty" function with PyTorch here.
I was wondering how I can implement this function in Tensorflow 2.
More specifically, I want to implement the function below, which is also in the code I shared its link.
def penalty(logits, y):
scale = torch.tensor(1.).cuda().requires_grad_()
loss = mean_nll(logits * scale, y)
grad = autograd.grad(loss, [scale], create_graph=True)[0]
return torch.sum(grad**2)