Merge pull request #174 from AnMakc/master

Improve inference throughput
This commit is contained in:
ShiYu
2025-11-05 17:48:13 +08:00
committed by GitHub
4 changed files with 74 additions and 74 deletions
+43 -17
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):
@@ -388,8 +388,6 @@ def sample_from_logits(logits, temperature=1.0, top_k=None, top_p=None, sample_l
def auto_regressive_inference(tokenizer, model, x, x_stamp, y_stamp, max_context, pred_len, clip=5, T=1.0, top_k=0, top_p=0.99, sample_count=5, verbose=False):
with torch.no_grad():
batch_size = x.size(0)
initial_seq_len = x.size(1)
x = torch.clip(x, -clip, clip)
device = x.device
@@ -398,14 +396,22 @@ def auto_regressive_inference(tokenizer, model, x, x_stamp, y_stamp, max_context
y_stamp = y_stamp.unsqueeze(1).repeat(1, sample_count, 1, 1).reshape(-1, y_stamp.size(1), y_stamp.size(2)).to(device)
x_token = tokenizer.encode(x, half=True)
initial_seq_len = x.size(1)
batch_size = x_token[0].size(0)
total_seq_len = initial_seq_len + pred_len
full_stamp = torch.cat([x_stamp, y_stamp], dim=1)
def get_dynamic_stamp(x_stamp, y_stamp, current_seq_len, pred_step):
generated_pre = x_token[0].new_empty(batch_size, pred_len)
generated_post = x_token[1].new_empty(batch_size, pred_len)
if current_seq_len <= max_context - pred_step:
return torch.cat([x_stamp, y_stamp[:, :pred_step, :]], dim=1)
else:
start_idx = max_context - pred_step
return torch.cat([x_stamp[:, -start_idx:, :], y_stamp[:, :pred_step, :]], dim=1)
pre_buffer = x_token[0].new_zeros(batch_size, max_context)
post_buffer = x_token[1].new_zeros(batch_size, max_context)
buffer_len = min(initial_seq_len, max_context)
if buffer_len > 0:
start_idx = max(0, initial_seq_len - max_context)
pre_buffer[:, :buffer_len] = x_token[0][:, start_idx:start_idx + buffer_len]
post_buffer[:, :buffer_len] = x_token[1][:, start_idx:start_idx + buffer_len]
if verbose:
ran = trange
@@ -413,13 +419,19 @@ def auto_regressive_inference(tokenizer, model, x, x_stamp, y_stamp, max_context
ran = range
for i in ran(pred_len):
current_seq_len = initial_seq_len + i
window_len = min(current_seq_len, max_context)
if current_seq_len <= max_context:
input_tokens = x_token
input_tokens = [
pre_buffer[:, :window_len],
post_buffer[:, :window_len]
]
else:
input_tokens = [t[:, -max_context:].contiguous() for t in x_token]
input_tokens = [pre_buffer, post_buffer]
current_stamp = get_dynamic_stamp(x_stamp, y_stamp, current_seq_len, i)
context_end = current_seq_len
context_start = max(0, context_end - max_context)
current_stamp = full_stamp[:, context_start:context_end, :].contiguous()
s1_logits, context = model.decode_s1(input_tokens[0], input_tokens[1], current_stamp)
s1_logits = s1_logits[:, -1, :]
@@ -429,14 +441,28 @@ def auto_regressive_inference(tokenizer, model, x, x_stamp, y_stamp, max_context
s2_logits = s2_logits[:, -1, :]
sample_post = sample_from_logits(s2_logits, temperature=T, top_k=top_k, top_p=top_p, sample_logits=True)
x_token[0] = torch.cat([x_token[0], sample_pre], dim=1)
x_token[1] = torch.cat([x_token[1], sample_post], dim=1)
generated_pre[:, i] = sample_pre.squeeze(-1)
generated_post[:, i] = sample_post.squeeze(-1)
torch.cuda.empty_cache()
if current_seq_len < max_context:
pre_buffer[:, current_seq_len] = sample_pre.squeeze(-1)
post_buffer[:, current_seq_len] = sample_post.squeeze(-1)
else:
pre_buffer.copy_(torch.roll(pre_buffer, shifts=-1, dims=1))
post_buffer.copy_(torch.roll(post_buffer, shifts=-1, dims=1))
pre_buffer[:, -1] = sample_pre.squeeze(-1)
post_buffer[:, -1] = sample_post.squeeze(-1)
input_tokens = [t[:, -max_context:].contiguous() for t in x_token]
full_pre = torch.cat([x_token[0], generated_pre], dim=1)
full_post = torch.cat([x_token[1], generated_post], dim=1)
context_start = max(0, total_seq_len - max_context)
input_tokens = [
full_pre[:, context_start:total_seq_len].contiguous(),
full_post[:, context_start:total_seq_len].contiguous()
]
z = tokenizer.decode(input_tokens, half=True)
z = z.reshape(batch_size, sample_count, z.size(1), z.size(2))
z = z.reshape(-1, sample_count, z.size(1), z.size(2))
preds = z.cpu().numpy()
preds = np.mean(preds, axis=1)
+16 -42
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:]
@@ -309,33 +312,6 @@ class RotaryPositionalEmbedding(nn.Module):
return torch.cat((-x2, x1), dim=-1)
def scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, scale=None, training=True) -> torch.Tensor:
L, S = query.size(-2), key.size(-2)
scale_factor = 1 / math.sqrt(query.size(-1)) if scale is None else scale
attn_bias = torch.zeros(L, S, dtype=query.dtype).to(query.device)
if is_causal:
assert attn_mask is None
temp_mask = torch.ones(L, S, dtype=torch.bool).tril(diagonal=0).to(query.device)
attn_bias.masked_fill_(temp_mask.logical_not(), float("-inf"))
attn_bias.to(query.dtype)
attn_weight = query @ key.transpose(-2, -1) * scale_factor
attn_weight += attn_bias
if attn_mask is not None:
attn_mask_bias = torch.zeros_like(attn_weight)
if attn_mask.dtype == torch.bool:
attn_mask_bias.masked_fill_(attn_mask, float("-inf"))
else:
attn_mask_bias += attn_mask
attn_weight += attn_mask_bias
attn_weight = torch.softmax(attn_weight, dim=-1)
attn_weight = torch.dropout(attn_weight, dropout_p, train=training)
return attn_weight @ value
class MultiHeadAttentionWithRoPE(nn.Module):
def __init__(self, d_model, n_heads, attn_dropout_p=0.0, resid_dropout_p=0.0):
super().__init__()
@@ -366,12 +342,11 @@ class MultiHeadAttentionWithRoPE(nn.Module):
else:
attn_mask = None
attn_output = scaled_dot_product_attention(
attn_output = F.scaled_dot_product_attention(
q, k, v,
attn_mask=attn_mask,
dropout_p=self.attn_dropout_p,
is_causal=True,
training=self.training
dropout_p=self.attn_dropout_p if self.training else 0.0,
is_causal=True
)
attn_output = attn_output.transpose(1, 2).contiguous().view(batch_size, seq_len, self.d_model)
@@ -411,12 +386,11 @@ class MultiHeadCrossAttentionWithRoPE(nn.Module):
is_causal_flag = self.training
attn_output = scaled_dot_product_attention(
attn_output = F.scaled_dot_product_attention(
q, k, v,
attn_mask=attn_mask,
dropout_p=self.attn_dropout_p,
is_causal=is_causal_flag,
training=self.training
dropout_p=self.attn_dropout_p if self.training else 0.0,
is_causal=is_causal_flag
)
attn_output = attn_output.transpose(1, 2).contiguous().view(batch_size, q_len, self.d_model)
+8 -8
View File
@@ -1,9 +1,9 @@
timestamps,open,high,low,close,volume,amount
2024-06-25 14:05:00,10.766402,10.778437,10.755835,10.769899,463.83276,479256.62
2024-06-25 14:10:00,10.769842,10.7804785,10.75896,10.771648,415.90912,434510.62
2024-06-25 14:15:00,10.771282,10.781633,10.760545,10.773098,396.62488,416206.88
2024-06-25 14:20:00,10.772831,10.782868,10.761984,10.77445,389.24976,409554.56
2024-06-25 14:25:00,10.774201,10.783865,10.763183,10.775418,386.3412,407075.44
2024-06-25 14:30:00,10.774968,10.78441,10.763903,10.776,383.4024,404050.8
2024-06-25 14:35:00,10.775348,10.7847595,10.764308,10.776471,377.25995,397440.12
2024-06-25 14:40:00,10.775859,10.78527,10.764823,10.77709,369.78687,389529.8
2024-06-25 14:05:00,10.766402,10.778437,10.755835,10.769899,463.83264,479256.5
2024-06-25 14:10:00,10.769841,10.7804785,10.75896,10.771648,415.90918,434510.62
2024-06-25 14:15:00,10.771282,10.781633,10.760545,10.773098,396.62482,416207.12
2024-06-25 14:20:00,10.772831,10.782868,10.761984,10.77445,389.2497,409554.62
2024-06-25 14:25:00,10.774201,10.783865,10.763183,10.775418,386.3413,407075.38
2024-06-25 14:30:00,10.774968,10.78441,10.763903,10.776,383.40247,404051.06
2024-06-25 14:35:00,10.775348,10.7847595,10.764308,10.776471,377.25995,397439.88
2024-06-25 14:40:00,10.775859,10.78527,10.764823,10.77709,369.787,389529.62
1 timestamps open high low close volume amount
2 2024-06-25 14:05:00 10.766402 10.778437 10.755835 10.769899 463.83276 463.83264 479256.62 479256.5
3 2024-06-25 14:10:00 10.769842 10.769841 10.7804785 10.75896 10.771648 415.90912 415.90918 434510.62
4 2024-06-25 14:15:00 10.771282 10.781633 10.760545 10.773098 396.62488 396.62482 416206.88 416207.12
5 2024-06-25 14:20:00 10.772831 10.782868 10.761984 10.77445 389.24976 389.2497 409554.56 409554.62
6 2024-06-25 14:25:00 10.774201 10.783865 10.763183 10.775418 386.3412 386.3413 407075.44 407075.38
7 2024-06-25 14:30:00 10.774968 10.78441 10.763903 10.776 383.4024 383.40247 404050.8 404051.06
8 2024-06-25 14:35:00 10.775348 10.7847595 10.764308 10.776471 377.25995 397440.12 397439.88
9 2024-06-25 14:40:00 10.775859 10.78527 10.764823 10.77709 369.78687 369.787 389529.8 389529.62
+7 -7
View File
@@ -1,9 +1,9 @@
timestamps,open,high,low,close,volume,amount
2024-07-03 09:55:00,10.897451,10.931036,10.800024,10.917972,1545.1384,1665960.5
2024-07-03 10:00:00,10.900613,10.907957,10.871778,10.884289,719.92456,792042.5
2024-07-03 09:55:00,10.897451,10.931036,10.800024,10.917972,1545.1383,1665960.5
2024-07-03 10:00:00,10.900613,10.907957,10.871778,10.884289,719.9246,792042.7
2024-07-03 10:05:00,10.882399,10.890674,10.864932,10.87375,659.0906,716546.6
2024-07-03 10:10:00,10.871227,10.881202,10.857913,10.867245,629.60645,681494.7
2024-07-03 10:15:00,10.864513,10.875556,10.85328,10.863286,607.7948,656790.1
2024-07-03 10:20:00,10.861447,10.872749,10.851165,10.86135,591.31,638401.9
2024-07-03 10:25:00,10.860088,10.871324,10.850175,10.860315,580.63446,626458.9
2024-07-03 10:30:00,10.858802,10.869965,10.849048,10.859139,572.116,616889.56
2024-07-03 10:10:00,10.871227,10.881202,10.857913,10.867245,629.6062,681494.6
2024-07-03 10:15:00,10.864513,10.875556,10.85328,10.863286,607.7946,656790.0
2024-07-03 10:20:00,10.861447,10.872749,10.851165,10.86135,591.30994,638401.56
2024-07-03 10:25:00,10.860088,10.871324,10.850175,10.860315,580.63446,626458.75
2024-07-03 10:30:00,10.858802,10.869965,10.849048,10.859139,572.1159,616889.56
1 timestamps open high low close volume amount
2 2024-07-03 09:55:00 10.897451 10.931036 10.800024 10.917972 1545.1384 1545.1383 1665960.5
3 2024-07-03 10:00:00 10.900613 10.907957 10.871778 10.884289 719.92456 719.9246 792042.5 792042.7
4 2024-07-03 10:05:00 10.882399 10.890674 10.864932 10.87375 659.0906 716546.6
5 2024-07-03 10:10:00 10.871227 10.881202 10.857913 10.867245 629.60645 629.6062 681494.7 681494.6
6 2024-07-03 10:15:00 10.864513 10.875556 10.85328 10.863286 607.7948 607.7946 656790.1 656790.0
7 2024-07-03 10:20:00 10.861447 10.872749 10.851165 10.86135 591.31 591.30994 638401.9 638401.56
8 2024-07-03 10:25:00 10.860088 10.871324 10.850175 10.860315 580.63446 626458.9 626458.75
9 2024-07-03 10:30:00 10.858802 10.869965 10.849048 10.859139 572.116 572.1159 616889.56