mirror of
https://github.com/shiyu-coder/Kronos.git
synced 2026-10-07 13:18:24 +08:00
Do not collect tokenizer metrics during inference
This commit is contained in:
+1
-1
@@ -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
@@ -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:]
|
||||
|
||||
Reference in New Issue
Block a user