mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
vectorized concat_latents_context cu_seq_k + ctx_encoder_type literal
This commit is contained in:
parent
b64abaceda
commit
033e027207
2 changed files with 19 additions and 22 deletions
|
|
@ -435,14 +435,14 @@ class CtxEncoderArguments:
|
|||
default=None,
|
||||
metadata={"help": "Context encoder model name or path."},
|
||||
)
|
||||
ctx_encoder_type: Literal[
|
||||
"embedding_only", "per_layer_activations", "early_exit"
|
||||
] = field(
|
||||
default="early_exit",
|
||||
metadata={
|
||||
"help": "Context encoder type. "
|
||||
"Options: 'embedding_only', 'per_layer_activations', 'early_exit'."
|
||||
},
|
||||
ctx_encoder_type: Literal["embed_only", "per_layer_activations", "early_exit"] = (
|
||||
field(
|
||||
default="early_exit",
|
||||
metadata={
|
||||
"help": "Context encoder type. "
|
||||
"Options: 'embed_only', 'per_layer_activations', 'early_exit'."
|
||||
},
|
||||
)
|
||||
)
|
||||
# used only with `early_exit` type
|
||||
layer_idx: int | None = field(
|
||||
|
|
|
|||
|
|
@ -385,9 +385,6 @@ class Idefics2PerceiverFlashAttention2(Idefics2PerceiverAttention):
|
|||
# slice context sample by sample and concat with latents
|
||||
old_cu_seq_lens_k = kwargs.pop("old_cu_seq_lens_k")
|
||||
cu_seq_lens_k = kwargs["cu_seq_lens_k"]
|
||||
old_cur_len = 0
|
||||
cur_len = 0
|
||||
n_latents = latents.shape[1]
|
||||
kv_inp = torch.empty(
|
||||
1,
|
||||
cu_seq_lens_k[-1],
|
||||
|
|
@ -395,18 +392,18 @@ class Idefics2PerceiverFlashAttention2(Idefics2PerceiverAttention):
|
|||
dtype=context.dtype,
|
||||
device=context.device,
|
||||
)
|
||||
for i, (old_cu_len, cu_len) in enumerate(
|
||||
zip(old_cu_seq_lens_k[1:], cu_seq_lens_k[1:])
|
||||
):
|
||||
ctx_len = old_cu_len - old_cur_len
|
||||
old_end_idx = cur_len + ctx_len
|
||||
kv_inp[0, cur_len:old_end_idx] = context[
|
||||
0, old_cur_len:old_cu_len
|
||||
# Compute context lengths and indices
|
||||
ctx_lens = old_cu_seq_lens_k[1:] - old_cu_seq_lens_k[:-1]
|
||||
ctx_start = old_cu_seq_lens_k[:-1]
|
||||
ctx_end = old_cu_seq_lens_k[1:]
|
||||
kv_start = cu_seq_lens_k[:-1]
|
||||
kv_end = cu_seq_lens_k[1:]
|
||||
# Fill context slices
|
||||
for i in range(bsz):
|
||||
kv_inp[0, kv_start[i] : kv_start[i] + ctx_lens[i]] = context[
|
||||
0, ctx_start[i] : ctx_end[i]
|
||||
]
|
||||
kv_inp[0, old_end_idx : old_end_idx + n_latents] = latents[i]
|
||||
|
||||
cur_len = cu_len
|
||||
old_cur_len = old_cu_len
|
||||
kv_inp[0, kv_start[i] + ctx_lens[i] : kv_end[i]] = latents[i]
|
||||
else:
|
||||
kv_inp = torch.cat([context, latents], dim=-2)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue