param_attr.py 2.2 KB
Newer Older
Y
Yu Yang 已提交
1 2 3
from initializer import Initializer, Xavier, Constant
from regularizer import WeightDecayRegularizer

Y
Yu Yang 已提交
4 5
__all__ = ['ParamAttr']

Y
Yu Yang 已提交
6 7 8 9 10 11 12

class ParamAttr(object):
    def __init__(self,
                 name=None,
                 initializer=None,
                 learning_rate=1.0,
                 regularizer=None,
Y
Yu Yang 已提交
13 14
                 trainable=True,
                 clip=None):
Y
Yu Yang 已提交
15 16 17 18 19
        self.name = name
        self.initializer = initializer
        self.learning_rate = learning_rate
        self.regularizer = regularizer
        self.trainable = trainable
Y
Yu Yang 已提交
20
        self.clip = clip
Y
Yu Yang 已提交
21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42

    def set_default_initializer(self, initializer):
        if initializer is None:
            if self.initializer is None:
                raise ValueError("ParamAttr.initializer is not set")
            return

        if self.initializer is not None:
            return

        self.initializer = initializer

    def set_default_param_initializer(self):
        self.set_default_initializer(Xavier())

    def set_default_bias_initializer(self):
        self.set_default_initializer(Constant(0.0))

    @staticmethod
    def to_attr(arg):
        if arg is None:
            return ParamAttr()
43 44
        elif isinstance(arg, list) or isinstance(arg, tuple):
            return [ParamAttr.to_attr(a) for a in arg]
Y
Yu Yang 已提交
45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62
        elif isinstance(arg, ParamAttr):
            return arg
        elif isinstance(arg, str) or isinstance(arg, unicode):
            return ParamAttr(name=arg)
        elif isinstance(arg, Initializer):
            return ParamAttr(initializer=arg)
        elif isinstance(arg, WeightDecayRegularizer):
            return ParamAttr(regularizer=arg)
        elif isinstance(arg, bool):
            return ParamAttr.to_attr(None) if arg else False
        else:
            raise TypeError("{0} cast to ParamAttr".format(type(arg)))

    def to_kwargs(self, with_initializer=False):
        kwargs = {
            'name': self.name,
            'learning_rate': self.learning_rate,
            'regularizer': self.regularizer,
Y
Yu Yang 已提交
63 64
            'trainable': self.trainable,
            'clip_attr': self.clip
Y
Yu Yang 已提交
65 66 67 68
        }
        if with_initializer:
            kwargs['initializer'] = self.initializer
        return kwargs