(self,
dense_column=None,
sparse_column=None,
mlp_bot=[512, 256, 64, 16],
mlp_top=[512, 256],
optimizer_type='adam',
learning_rate=0.1,
inputs=None,
interaction_op='dot',
bf16=False,
stock_tf=None,
adaptive_emb=False,
input_layer_partitioner=None,
dense_layer_partitioner=None)
| 67 | |
| 68 | class DLRM(): |
| 69 | def __init__(self, |
| 70 | dense_column=None, |
| 71 | sparse_column=None, |
| 72 | mlp_bot=[512, 256, 64, 16], |
| 73 | mlp_top=[512, 256], |
| 74 | optimizer_type='adam', |
| 75 | learning_rate=0.1, |
| 76 | inputs=None, |
| 77 | interaction_op='dot', |
| 78 | bf16=False, |
| 79 | stock_tf=None, |
| 80 | adaptive_emb=False, |
| 81 | input_layer_partitioner=None, |
| 82 | dense_layer_partitioner=None): |
| 83 | if not inputs: |
| 84 | raise ValueError('Dataset is not defined.') |
| 85 | self._feature = inputs[0] |
| 86 | self._label = inputs[1] |
| 87 | |
| 88 | if not dense_column or not sparse_column: |
| 89 | raise ValueError('Dense column or sparse column is not defined.') |
| 90 | self._dense_column = dense_column |
| 91 | self._sparse_column = sparse_column |
| 92 | |
| 93 | self.tf = stock_tf |
| 94 | self.bf16 = False if self.tf else bf16 |
| 95 | self.is_training = True |
| 96 | self._adaptive_emb = adaptive_emb |
| 97 | |
| 98 | self._mlp_bot = mlp_bot |
| 99 | self._mlp_top = mlp_top |
| 100 | self._learning_rate = learning_rate |
| 101 | self._input_layer_partitioner = input_layer_partitioner |
| 102 | self._dense_layer_partitioner = dense_layer_partitioner |
| 103 | self._optimizer_type = optimizer_type |
| 104 | self.interaction_op = interaction_op |
| 105 | if self.interaction_op not in ['dot', 'cat']: |
| 106 | print("Invaild interaction op, must be 'dot' or 'cat'.") |
| 107 | sys.exit() |
| 108 | |
| 109 | self._create_model() |
| 110 | with tf.name_scope('head'): |
| 111 | self._create_loss() |
| 112 | self._create_optimizer() |
| 113 | self._create_metrics() |
| 114 | |
| 115 | # used to add summary in tensorboard |
| 116 | def _add_layer_summary(self, value, tag): |
nothing calls this directly
no test coverage detected