Is there an equivalent of _keras_shape in pytorch?

Viewed 150

I am trying to convert some code from a keras implementation to a pytorch equivalent. I don't understand this assignment in particular f_real._keras_shape = self.kernel_shape. As I can see the _keras_shape seems to be an auto-generated attribute of 'self.kernel' - which is being assigned the kernel_shape value. In keras, I think the kernel is initialized as a tensor placeholder of sorts. Here's the code:

from keras.layers import Layer
from keras import backend as K
class cconv(Layer):
.......
  self.kernel = self.add_weight(
        self.kernel_shape,
        initializer=kern_init,
        name='kernel',
        regularizer=self.kernel_regularizer,
        constraint=self.kernel_constraint
    )
  real = self.kernel[:, :, :, :self.filters]
  imag = self.kernel[:, :, :, self.filters:]

I am struggling with these 2 lines:

  real._keras_shape = self.kernel_shape
  imag._keras_shape = self.kernel_shape

so far, I got:

  self.kernel = nn.Parameter(
        self.kernel_shape,
        initializer=self.kernel_initializer,
        name='kernel',
        regularizer=self.kernel_regularizer,
        constraint=self.kernel_constraint
    )
  real = self.kernel[:, :, :, :self.filters]
  imag = self.kernel[:, :, :, self.filters:]

Is there a torch.nn.Parameter equivalent of '_keras_shape' or any workaround to this?

edit: I have done some digging but I can't seem to find the exact file where this '_keras_shape' attribute originates! There is a Variable class which seems relevant (can't reach the bottom), i.e. some code in keras.backend.py -

from tensorflow.python.ops import variables as variables_module

def variable(value, dtype=None, name=None, constraint=None):

    v = variables_module.Variable(
        value,
        dtype=dtypes_module.as_dtype(dtype),
        name=name,
        constraint=constraint)
    if isinstance(value, np.ndarray):
        v._keras_shape = value.shape
    elif hasattr(value, 'shape'):
        v._keras_shape = int_shape(value)
    track_variable(v)
    return v

The variables.py file does not make sense - can't find how far this _keras_shape goes back to...

0 Answers
Related