Do not collect tokenizer metrics during inference

This commit is contained in:
Maxim
2025-10-29 12:04:43 -03:00
parent dccfa764fc
commit 21eb36afc7
2 changed files with 11 additions and 8 deletions
+1 -1
View File
@@ -155,7 +155,7 @@ class KronosTokenizer(nn.Module, PyTorchModelHubMixin):
z = layer(z)
z = self.quant_embed(z)
bsq_loss, quantized, z_indices = self.tokenizer(z, half)
bsq_loss, quantized, z_indices = self.tokenizer(z, half=half, collect_metrics=False)
return z_indices
def decode(self, x, half=False):
+10 -7
View File
@@ -87,11 +87,18 @@ class BinarySphericalQuantizer(nn.Module):
torch.tensor(-1, dtype=z.dtype, device=z.device))
return z + (zhat - z).detach()
def forward(self, z):
def forward(self, z, collect_metrics=True):
# if self.input_format == 'bchw':
# z = rearrange(z, 'b c h w -> b h w c')
zq = self.quantize(z)
q_scale = 1. / (self.embed_dim ** 0.5) if self.l2_norm else 1.
zq = zq * q_scale
if not collect_metrics:
return zq, zq.new_zeros(()), {}
indices = self.codes_to_indexes(zq.detach())
group_indices = self.codes_to_group_indexes(zq.detach())
if not self.training:
@@ -99,8 +106,6 @@ class BinarySphericalQuantizer(nn.Module):
else:
used_codes = None
q_scale = 1. / (self.embed_dim ** 0.5) if self.l2_norm else 1.
if self.soft_entropy:
persample_entropy, cb_entropy, avg_prob = self.soft_entropy_loss(z)
entropy_penalty = self.gamma0 * persample_entropy - self.gamma * cb_entropy
@@ -110,8 +115,6 @@ class BinarySphericalQuantizer(nn.Module):
cb_entropy = codebook_entropy(zq, self.basis, self.embed_dim)
entropy_penalty = self.gamma0 * persample_entropy - self.gamma * cb_entropy
zq = zq * q_scale
# commit loss
commit_loss = self.beta * torch.mean(((zq.detach() - z) ** 2).sum(dim=-1))
@@ -239,9 +242,9 @@ class BSQuantizer(nn.Module):
)
return (bits * indices).sum(-1)
def forward(self, z, half=False):
def forward(self, z, half=False, collect_metrics=True):
z = F.normalize(z, dim=-1)
quantized, bsq_loss, metrics = self.bsq(z)
quantized, bsq_loss, metrics = self.bsq(z, collect_metrics=collect_metrics)
if half:
q_pre = quantized[:, :, :self.s1_bits]
q_post = quantized[:, :, self.s1_bits:]