Add T5LA models - #39293
Conversation
| ("swinv2", "Swin Transformer V2"), | ||
| ("switch_transformers", "SwitchTransformers"), | ||
| ("t5", "T5"), | ||
| ("t5la", "T5LALA"), |
There was a problem hiding this comment.
Why add-new-model-like has added two "LA" suffixes? It has to be "T5LA" I guess. I double checked my answers to that command and even retried it a second time. Everything was correct but this output doesn't look right!
3aae8b0 to
06d9258
Compare
ArthurZucker
left a comment
There was a problem hiding this comment.
Thanks for adding the model! 🚀
| ("switch_transformers", "SwitchTransformers"), | ||
| ("t5", "T5"), | ||
| ("t5gemma", "T5Gemma"), | ||
| ("t5la", "T5LALA"), |
There was a problem hiding this comment.
A bit wird to have this no?
| from collections.abc import Mapping | ||
|
|
||
| from ...configuration_utils import PretrainedConfig | ||
| from ...onnx import OnnxSeq2SeqConfigWithPast |
There was a problem hiding this comment.
no longer used for new models!
| """ | ||
|
|
||
| model_type = "t5la" | ||
| lookahead_type = "la" # Other options are "laa", "laa2", and "lae" |
There was a problem hiding this comment.
think we should add this as an arg no?
| # for backwards compatibility | ||
| if feed_forward_proj == "gated-gelu": | ||
| self.dense_act_fn = "gelu_new" |
There was a problem hiding this comment.
new model don't have BC
| class T5LAOnnxConfig(OnnxSeq2SeqConfigWithPast): | ||
| @property | ||
| def inputs(self) -> Mapping[str, Mapping[int, str]]: | ||
| common_inputs = { | ||
| "input_ids": {0: "batch", 1: "encoder_sequence"}, | ||
| "attention_mask": {0: "batch", 1: "encoder_sequence"}, | ||
| } | ||
| if self.use_past: | ||
| common_inputs["attention_mask"][1] = "past_encoder_sequence + sequence" | ||
| common_inputs["decoder_input_ids"] = {0: "batch"} | ||
| common_inputs["decoder_attention_mask"] = {0: "batch", 1: "past_decoder_sequence + sequence"} | ||
| else: | ||
| common_inputs["decoder_input_ids"] = {0: "batch", 1: "decoder_sequence"} | ||
| common_inputs["decoder_attention_mask"] = {0: "batch", 1: "decoder_sequence"} | ||
|
|
||
| if self.use_past: | ||
| self.fill_with_past_key_values_(common_inputs, direction="inputs") | ||
|
|
||
| return common_inputs | ||
|
|
||
| @property | ||
| def default_onnx_opset(self) -> int: | ||
| return 13 | ||
|
|
There was a problem hiding this comment.
| class T5LAOnnxConfig(OnnxSeq2SeqConfigWithPast): | |
| @property | |
| def inputs(self) -> Mapping[str, Mapping[int, str]]: | |
| common_inputs = { | |
| "input_ids": {0: "batch", 1: "encoder_sequence"}, | |
| "attention_mask": {0: "batch", 1: "encoder_sequence"}, | |
| } | |
| if self.use_past: | |
| common_inputs["attention_mask"][1] = "past_encoder_sequence + sequence" | |
| common_inputs["decoder_input_ids"] = {0: "batch"} | |
| common_inputs["decoder_attention_mask"] = {0: "batch", 1: "past_decoder_sequence + sequence"} | |
| else: | |
| common_inputs["decoder_input_ids"] = {0: "batch", 1: "decoder_sequence"} | |
| common_inputs["decoder_attention_mask"] = {0: "batch", 1: "decoder_sequence"} | |
| if self.use_past: | |
| self.fill_with_past_key_values_(common_inputs, direction="inputs") | |
| return common_inputs | |
| @property | |
| def default_onnx_opset(self) -> int: | |
| return 13 |
| ) | ||
|
|
||
| if split_mlp_wi: | ||
| flax_model.params["decoder"]["block"][str(layer_index)]["layer"]["2"]["DenseReluDense"]["wi_0"][ |
There was a problem hiding this comment.
| flax_model.params["decoder"]["block"][str(layer_index)]["layer"]["2"]["DenseReluDense"]["wi_0"][ | |
| flax_model.params["decoder"]["block"][str(layer_index)]["layer"]["2"]["DenseReluDense"]["wi_0"][ |
we should define a utils like a def flax_get_flatten to which you give a key and it just does this!
| k, o, q, v = t5lax_attention_lookup(old, i, "encoder", "attention") | ||
| new[f"encoder.block.{i}.layer.0.layer_norm.weight"] = layer_norm | ||
| new[f"encoder.block.{i}.layer.0.SelfAttention.k.weight"] = k.T | ||
| new[f"encoder.block.{i}.layer.0.SelfAttention.o.weight"] = o.T | ||
| new[f"encoder.block.{i}.layer.0.SelfAttention.q.weight"] = q.T | ||
| new[f"encoder.block.{i}.layer.0.SelfAttention.v.weight"] = v.T | ||
|
|
||
| # Block i, layer 1 (MLP). | ||
| layer_norm = t5lax_layer_norm_lookup(old, i, "encoder", "pre_mlp_layer_norm") | ||
| wi, wo = t5lax_mlp_lookup(old, i, "encoder", split_mlp_wi) | ||
| new[f"encoder.block.{i}.layer.1.layer_norm.weight"] = layer_norm | ||
| if split_mlp_wi: | ||
| new[f"encoder.block.{i}.layer.1.DenseReluDense.wi_0.weight"] = wi[0].T | ||
| new[f"encoder.block.{i}.layer.1.DenseReluDense.wi_1.weight"] = wi[1].T | ||
| else: | ||
| new[f"encoder.block.{i}.layer.1.DenseReluDense.wi.weight"] = wi.T | ||
| new[f"encoder.block.{i}.layer.1.DenseReluDense.wo.weight"] = wo.T | ||
|
|
||
| new["encoder.block.0.layer.0.SelfAttention.relative_attention_bias.weight"] = old[ | ||
| "encoder/relpos_bias/rel_embedding" | ||
| ].T | ||
| new["encoder.final_layer_norm.weight"] = old["encoder/encoder_norm/scale"] | ||
|
|
||
| if not is_encoder_only: | ||
| # Decoder. | ||
| for i in range(num_decoder_layers): | ||
| # Block i, layer 0 (Self Attention). | ||
| layer_norm = t5lax_layer_norm_lookup(old, i, "decoder", "pre_self_attention_layer_norm") | ||
| k, o, q, v = t5lax_attention_lookup(old, i, "decoder", "self_attention") | ||
| new[f"decoder.block.{i}.layer.0.layer_norm.weight"] = layer_norm | ||
| new[f"decoder.block.{i}.layer.0.SelfAttention.k.weight"] = k.T | ||
| new[f"decoder.block.{i}.layer.0.SelfAttention.o.weight"] = o.T | ||
| new[f"decoder.block.{i}.layer.0.SelfAttention.q.weight"] = q.T | ||
| new[f"decoder.block.{i}.layer.0.SelfAttention.v.weight"] = v.T | ||
|
|
||
| # Block i, layer 1 (Cross Attention). | ||
| layer_norm = t5lax_layer_norm_lookup(old, i, "decoder", "pre_cross_attention_layer_norm") | ||
| k, o, q, v = t5lax_attention_lookup(old, i, "decoder", "encoder_decoder_attention") | ||
| new[f"decoder.block.{i}.layer.1.layer_norm.weight"] = layer_norm | ||
| new[f"decoder.block.{i}.layer.1.EncDecAttention.k.weight"] = k.T | ||
| new[f"decoder.block.{i}.layer.1.EncDecAttention.o.weight"] = o.T | ||
| new[f"decoder.block.{i}.layer.1.EncDecAttention.q.weight"] = q.T | ||
| new[f"decoder.block.{i}.layer.1.EncDecAttention.v.weight"] = v.T | ||
|
|
||
| # Block i, layer 2 (MLP). | ||
| layer_norm = t5lax_layer_norm_lookup(old, i, "decoder", "pre_mlp_layer_norm") | ||
| wi, wo = t5lax_mlp_lookup(old, i, "decoder", split_mlp_wi) | ||
| new[f"decoder.block.{i}.layer.2.layer_norm.weight"] = layer_norm | ||
| if split_mlp_wi: | ||
| new[f"decoder.block.{i}.layer.2.DenseReluDense.wi_0.weight"] = wi[0].T | ||
| new[f"decoder.block.{i}.layer.2.DenseReluDense.wi_1.weight"] = wi[1].T | ||
| else: | ||
| new[f"decoder.block.{i}.layer.2.DenseReluDense.wi.weight"] = wi.T | ||
| new[f"decoder.block.{i}.layer.2.DenseReluDense.wo.weight"] = wo.T | ||
|
|
||
| new["decoder.final_layer_norm.weight"] = old["decoder/decoder_norm/scale"] | ||
| new["decoder.block.0.layer.0.SelfAttention.relative_attention_bias.weight"] = old[ | ||
| "decoder/relpos_bias/rel_embedding" | ||
| ].T | ||
|
|
||
| # LM Head (only in v1.1 checkpoints, in v1.0 embeddings are used instead) |
There was a problem hiding this comment.
take inspiration from the convert_llama4 file for an update to this!
|
|
||
| import numpy as np | ||
| import tensorflow as tf | ||
| except ImportError: | ||
| logger.error( | ||
| "Loading a TensorFlow model in PyTorch, requires TensorFlow to be installed. Please see " | ||
| "https://www.tensorflow.org/install/ for installation instructions." | ||
| ) | ||
| raise | ||
| tf_path = os.path.abspath(tf_checkpoint_path) | ||
| logger.info(f"Converting TensorFlow checkpoint from {tf_path}") | ||
| # Load weights from TF model | ||
| init_vars = tf.train.list_variables(tf_path) | ||
| names = [] | ||
| tf_weights = {} | ||
| for name, shape in init_vars: | ||
| logger.info(f"Loading TF weight {name} with shape {shape}") | ||
| array = tf.train.load_variable(tf_path, name) | ||
| names.append(name) | ||
| tf_weights[name] = array | ||
|
|
||
| for txt_name in names: | ||
| name = txt_name.split("/") | ||
| # adam_v and adam_m are variables used in AdamWeightDecayOptimizer to calculated m and v | ||
| # which are not required for using pretrained model | ||
| if any( | ||
| n in ["adam_v", "adam_m", "AdamWeightDecayOptimizer", "AdamWeightDecayOptimizer_1", "global_step"] | ||
| for n in name | ||
| ): | ||
| logger.info(f"Skipping {'/'.join(name)}") | ||
| tf_weights.pop(txt_name, None) | ||
| continue | ||
| if "_slot_" in name[-1]: | ||
| logger.info(f"Skipping {'/'.join(name)}") | ||
| tf_weights.pop(txt_name, None) | ||
| continue | ||
| pointer = model | ||
| array = tf_weights[txt_name] | ||
|
|
||
| for m_name in name: | ||
| if re.fullmatch(r"[A-Za-z]+_\d+", m_name): | ||
| scope_names = re.split(r"_(\d+)", m_name) | ||
| else: | ||
| scope_names = [m_name] | ||
| if scope_names[0] in ["kernel", "scale", "embedding"]: | ||
| pointer = getattr(pointer, "weight") | ||
| elif scope_names[0] == "self_attention": | ||
| pointer = getattr(pointer, "layer") | ||
| pointer = pointer[0] | ||
| elif scope_names[0] == "enc_dec_attention": | ||
| pointer = getattr(pointer, "layer") | ||
| pointer = pointer[1] | ||
| elif scope_names[0] == "dense_relu_dense": | ||
| pointer = getattr(pointer, "layer") | ||
| pointer = pointer[2] | ||
| elif scope_names[0] == "rms_norm": | ||
| if hasattr(pointer, "layer_norm"): | ||
| pointer = getattr(pointer, "layer_norm") | ||
| elif hasattr(pointer, "final_layer_norm"): | ||
| pointer = getattr(pointer, "final_layer_norm") | ||
| elif scope_names[0] == "scale": | ||
| pointer = getattr(pointer, "weight") | ||
| elif scope_names[0] == "output_bias" or scope_names[0] == "beta": | ||
| pointer = getattr(pointer, "bias") | ||
| elif scope_names[0] == "squad": | ||
| pointer = getattr(pointer, "classifier") | ||
| elif scope_names[0] == "decoder" and name[1] == "logits": | ||
| continue | ||
| elif scope_names[0] == "logits": | ||
| pointer = getattr(pointer, "lm_head") | ||
| elif scope_names[0] == "wi" and len(scope_names) > 1 and scope_names[1].isdigit(): | ||
| pointer = getattr(pointer, f"wi_{scope_names[1]}") | ||
| continue | ||
| else: | ||
| try: | ||
| pointer = getattr(pointer, scope_names[0]) | ||
| except AttributeError: | ||
| logger.info(f"Skipping {'/'.join(name)}") | ||
| continue | ||
| if len(scope_names) >= 2: | ||
| num = int(scope_names[1]) | ||
| pointer = pointer[num] | ||
| if scope_names[0] not in ["kernel", "scale", "embedding"]: | ||
| pointer = getattr(pointer, "weight") | ||
| if scope_names[0] != "embedding": | ||
| logger.info(f"Transposing numpy weight of shape {array.shape} for {name}") | ||
| array = np.transpose(array) | ||
| try: | ||
| if pointer.shape != array.shape: | ||
| raise ValueError(f"Pointer shape {pointer.shape} and array shape {array.shape} mismatched") | ||
| except AssertionError as e: | ||
| e.args += (pointer.shape, array.shape) | ||
| raise | ||
| logger.info(f"Initialize PyTorch weight {name}") | ||
| pointer.data = torch.from_numpy(array.astype(np.float32)) | ||
| tf_weights.pop(txt_name, None) | ||
|
|
||
| logger.info(f"Weights not copied to PyTorch model: {', '.join(tf_weights.keys())}.") | ||
| return model | ||
|
|
There was a problem hiding this comment.
t5 model is very old, recommending you to take inspiration from T5Gemma that was just added, and use modular to isolate the differences! 🤗
There was a problem hiding this comment.
Thank you for your feedback @ArthurZucker, and also for introducing modular approach and the Gemma, I wasn't familiar with them!
I created a separate branch for my copy of gemma, available here, would you recommend that I force push it here on this branch and PR? Or should I close this and create a new PR?
There was a problem hiding this comment.
You can force push here will be simpler 🫡
There was a problem hiding this comment.
Sure, hopefully soon.
By the way, as Gemma's structure is relatively different from T5, when you said "t5 model is very old", did you mean in terms of architecture and design or in terms of the way it is implemented in Transformers (modular approach)? Whilst I'm happy to provide the LA version of Gemma (as is in progress in the mentioned branch), I still would need a model as close to T5 as possible, for the purpose of comparison in my paper (I guess). So what's your advice in this regard?
There was a problem hiding this comment.
The best example is T5Gemma as is adheres to the recent standards of the transformers library!
There was a problem hiding this comment.
The idea is just that our code in transformers is old
so has a lot of legacy
ca104a5 to
1b7338e
Compare
1b7338e to
c0aceaa
Compare
|
[For maintainers] Suggested jobs to run (before merge) run-slow: auto, t5la |
What does this PR do?
Adds the implementation of the LookAhead (LA) models. These models are designed to predict not only the next immediate token after the input prompt, but also the second, third, ... up to K next tokens. These tokens can be used to mitigate the high inference latency in generation (see this survey) or in approximated ranking of a set of responses (see this paper for an application).
Before submitting
Pull Request section?
to it if that's the case.
documentation guidelines, and
here are tips on formatting docstrings.
Who can review?
@ArthurZucker