Skip to content

Instantly share code, notes, and snippets.

for key in model.fc.state_dict():
print('key: ', key)
param = model.fc.state_dict()[key]
print('param.shape: ', param.shape)
print('param.requires_grad: ', param.requires_grad)
print('param.shape, param.requires_grad: ', param.shape, param.requires_grad)
print('isinstance(param, nn.Module) ', isinstance(param, nn.Module))
print('isinstance(param, nn.Parameter) ', isinstance(param, nn.Parameter))
print('isinstance(param, torch.Tensor): ', isinstance(param, torch.Tensor))
print('=====')
torch.save(model.state_dict(), 'weights_only.pth')
model_new = NeuralNet()
model_new.load_state_dict(torch.load('weights_only.pth'))
for name, param in model_new.named_parameters():
print(name, ':', param.requires_grad)
torch.save(model, 'entire_model.pth')
model_new = torch.load('entire_model.pth')
for name, param in model_new.named_parameters():
print(name, ':', param.requires_grad)
for name , param in model.named_parameters():
print('type(param): ', type(param))
print('isinstance(param, nn.Module): ', isinstance(param, nn.Module))
print('isinstance(param, nn.Parameter): ', isinstance(param, nn.Parameter))
print('isinstance(param, torch.Tensor) ', isinstance(param, torch.Tensor))
print('=====')
# Also try model.parameters(). It doesn't return the name of the parameters but just the parameters.
for param in model.layer1.parameters():
param.requires_grad = False
for name, child in model.named_children():
print('name: ', name)
print('isinstance(child, nn.Module): ', isinstance(child, nn.Module))
print('isinstance(child, nn.Parameter): ', isinstance(child, nn.Parameter))
print('isinstance(child, torch.Tensor) ', isinstance(child, torch.Tensor))
print('=====')
# Also try model.children(). It doesn't return the name of the children, but the children (nn.Module objects)