Complex data to neural network in PyTorch 1.8.1

Viewed 440

I am trying to used complex valued data as input to a test neural network. From the release notes (point 2), PyTorch 1.8.0 is said to support complex autograd. My code that I used to test this functionality is as follows. I am using 1.8.1 version of the library.

import torch
from torch import nn, optim


class ComplexTest(nn.Module):
    def __init__(self):
        super(ComplexTest, self).__init__()
        self.fc1 = nn.Linear(10, 20)
        self.fc2 = nn.Linear(20, 10)
        self.relu = nn.ReLU()

    def forward(self, inputs):
        return self.fc2(self.relu(self.fc1(inputs)))


device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
complex_test = ComplexTest().to(device)
complex_test.train()

opt = optim.Adam(complex_test.parameters())

mse_loss = nn.MSELoss()

for _ in range(100):
    opt.zero_grad()

    inp = torch.randn((1000, 10), dtype=torch.cfloat).to(device)
    op = complex_test(inp)

    loss = mse_loss(op, inp)
    loss.backward()
    opt.step()
    print(loss.item())

But this gives an error

expected scalar type Float but found ComplexFloat

Is this not supported or did I read the documentation wrong? Thanks!

0 Answers
Related