forked from IIT-Lab/complexPyTorch
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathcomplexFunctions.py
More file actions
28 lines (20 loc) · 989 Bytes
/
Copy pathcomplexFunctions.py
File metadata and controls
28 lines (20 loc) · 989 Bytes
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
@author: spopoff
"""
from torch.nn.functional import relu, max_pool2d, dropout, dropout2d
def complex_relu(input_r,input_i):
return relu(input_r), relu(input_i)
def complex_max_pool2d(input_r,input_i,kernel_size, stride=None, padding=0,
dilation=1, ceil_mode=False, return_indices=False):
return max_pool2d(input_r, kernel_size, stride, padding, dilation,
ceil_mode, return_indices), \
max_pool2d(input_i, kernel_size, stride, padding, dilation,
ceil_mode, return_indices)
def complex_dropout(input_r,input_i, p=0.5, training=True, inplace=False):
return dropout(input_r, p, training, inplace), \
dropout(input_i, p, training, inplace)
def complex_dropout2d(input_r,input_i, p=0.5, training=True, inplace=False):
return dropout2d(input_r, p, training, inplace), \
dropout2d(input_i, p, training, inplace)