inputs: (batch_size, seq_len, emb_dim)
(self, inputs)
| 627 | return text_enc_out, text_seq_len, text_self_attn_mask, checkpoints |
| 628 | |
| 629 | def vector_quantizer(self, inputs): |
| 630 | """inputs: (batch_size, seq_len, emb_dim)""" |
| 631 | input_shape = paddle.shape(inputs) |
| 632 | # Flatten input (batch_size * seq_len, emb_dim) |
| 633 | flat_input = paddle.reshape(inputs, shape=[-1, self._emb_size]) |
| 634 | |
| 635 | # Calculate distances (batch_size * seq_len, num_codebook) |
| 636 | distances = (paddle.sum(paddle.pow(flat_input, 2), axis=1, keepdim=True) |
| 637 | + paddle.unsqueeze(x=paddle.sum(paddle.pow(self.vq_emb, 2), axis=1), axis=0) |
| 638 | - 2 * paddle.matmul(flat_input, paddle.transpose(self.vq_emb, perm=[1, 0]))) |
| 639 | |
| 640 | # Encoding (batch_size * seq_len, 1) |
| 641 | encoding_indices = paddle.unsqueeze(x=paddle.argmin(distances, axis=1), axis=1) |
| 642 | # paddle.static.Print(encoding_indices, message="encoding_indices", summarize=1000) |
| 643 | size_range = paddle.unsqueeze(x=paddle.arange(paddle.shape(encoding_indices)[0], dtype='int64'), axis=1) |
| 644 | # (batch_size * seq_len, 2) |
| 645 | index = paddle.concat([size_range, encoding_indices], axis=1) |
| 646 | |
| 647 | # (batch_size * seq_len, num_codebook) |
| 648 | out_shape = [paddle.shape(encoding_indices)[0], self.num_codebook] |
| 649 | # (batch_size * seq_len) |
| 650 | updates = paddle.ones(shape=[paddle.shape(encoding_indices)[0]], dtype='float32') |
| 651 | # (batch_size * seq_len, num_codebook) |
| 652 | encodings = paddle.scatter_nd(index=index, updates=updates, shape=out_shape) |
| 653 | # paddle.static.Print(encodings, message="encodings", summarize=1000) |
| 654 | |
| 655 | # Quantize and unflatten (batch_size, seq_len, emb_dim) |
| 656 | quantized = paddle.reshape(x=paddle.matmul(encodings, self.vq_emb), |
| 657 | shape=[input_shape[0], input_shape[1], self._emb_size]) |
| 658 | # paddle.static.Print(quantized, message="quantized", summarize=1000) |
| 659 | encoding_indices_reshaped = paddle.reshape(encoding_indices, shape=[input_shape[0], input_shape[1]]) |
| 660 | |
| 661 | constant_zeros = paddle.zeros_like(x=inputs) |
| 662 | quantized_detach = paddle.add(constant_zeros, quantized) |
| 663 | quantized_detach.stop_gradient = True |
| 664 | inputs_detach = paddle.add(constant_zeros, inputs) |
| 665 | inputs_detach.stop_gradient = True |
| 666 | |
| 667 | # Loss |
| 668 | e_latent_loss = paddle.nn.functional.mse_loss(input=quantized_detach, label=inputs) |
| 669 | # paddle.static.Print(e_latent_loss, message="e_latent_loss", summarize=1000) |
| 670 | q_latent_loss = paddle.nn.functional.mse_loss(input=quantized, label=inputs_detach) |
| 671 | # paddle.static.Print(q_latent_loss, message="q_latent_loss", summarize=1000) |
| 672 | loss = self._alpha * q_latent_loss + self._beta * e_latent_loss |
| 673 | |
| 674 | # Straight Through Estimator |
| 675 | # sg_quantized = paddle.subtract(x=quantized, y=inputs) |
| 676 | # sg_quantized.stop_gradient = True |
| 677 | # quantized_emb = paddle.add(x=inputs, y=sg_quantized) |
| 678 | # paddle.static.Print(quantized_emb, message="quantized_emb", summarize=1000) |
| 679 | return loss, quantized, encoding_indices_reshaped |
| 680 | |
| 681 | def topk_vector_quantizer(self, inputs, K=100): |
| 682 | """inputs: (batch_size, seq_len, emb_dim)""" |
no test coverage detected