MCPcopy Create free account
hub / github.com/apple/axlearn / forward

Method forward

axlearn/common/quantizer.py:362–407  ·  view source on GitHub ↗

Computes random projection and quantization. Args: inputs: Tensor of shape [batch_size, seq_len, input_dim]. paddings: 0/1 Tensor of shape [batch_size, seq_len]. Returns: BaseQuantizer.Output.

(self, inputs: Tensor, *, paddings: Tensor)

Source from the content-addressed store, hash-verified

360 return params
361
362 def forward(self, inputs: Tensor, *, paddings: Tensor) -> BaseQuantizer.Output:
363 """Computes random projection and quantization.
364
365 Args:
366 inputs: Tensor of shape [batch_size, seq_len, input_dim].
367 paddings: 0/1 Tensor of shape [batch_size, seq_len].
368
369 Returns:
370 BaseQuantizer.Output.
371 """
372 cfg = self.config
373
374 # [batch_size, seq_len, num_codebooks * codebook_dim].
375 inputs = self.rand_proj(inputs)
376 inputs_by_group = jnp.reshape(
377 inputs, list(inputs.shape[:2]) + [cfg.num_codebooks, cfg.codebook_dim]
378 )
379
380 if cfg.normalize_inputs:
381 # [..., num_codebooks, codebook_dim].
382 inputs_by_group = l2_normalize(inputs_by_group, axis=-1, eps=1e-12)
383
384 # When codebook is normalized, dot_product is equivalent to l2_distance.
385 metric = (
386 SimilarityMetric.DOT_PRODUCT if cfg.normalize_codebook else SimilarityMetric.L2_DISTANCE
387 )
388 q_outputs = quantize_by_nearest_neighbor(
389 inputs=inputs_by_group,
390 codebook=self.parameters["codebook"],
391 metric=metric,
392 )
393 q_outputs = _apply_paddings(outputs=q_outputs, paddings=paddings)
394 # Best-rq freezes the codebook.
395 ids = jax.lax.stop_gradient(q_outputs.ids)
396 quantized_vectors = jax.lax.stop_gradient(q_outputs.quantized_vectors)
397
398 outputs = self.Output(
399 # [batch_size, seq_len, num_codebooks].
400 ids=ids,
401 # [batch_size, seq_len, num_codebooks, codebook_dim].
402 quantized_vectors=quantized_vectors,
403 )
404
405 onehots = _ids_to_onehots(outputs.ids, codebook_size=cfg.codebook_size, dtype=jnp.int32)
406 _add_codebook_summaries(context=current_context(), onehots=onehots, paddings=paddings)
407 return outputs
408
409
410class KmeansVectorQuantizer(BaseQuantizer):

Callers

nothing calls this directly

Calls 6

l2_normalizeFunction · 0.90
current_contextFunction · 0.90
_apply_paddingsFunction · 0.85
_ids_to_onehotsFunction · 0.85
_add_codebook_summariesFunction · 0.85

Tested by

no test coverage detected