MCPcopy Create free account
hub / github.com/PaddlePaddle/Research / vector_quantizer

Method vector_quantizer

NLP/UNIMO-2/src/model/unimo_grounded.py:629–679  ·  view source on GitHub ↗

inputs: (batch_size, seq_len, emb_dim)

(self, inputs)

Source from the content-addressed store, hash-verified

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)"""

Callers 1

_gen_inputMethod · 0.95

Calls 2

transposeMethod · 0.45
addMethod · 0.45

Tested by

no test coverage detected