From 0abe15875ce990403a5c2c3c9b09954d36e99a65 Mon Sep 17 00:00:00 2001 From: Dobiasd Date: Sun, 3 May 2026 18:48:07 +0200 Subject: [PATCH] Add use_causal_mask to MultiHeadAttention and approximate GELU The MultiHeadAttention layer now respects the per-call use_causal_mask=True kwarg: scores at (t, k) for k > t are set to -inf before softmax. The converter extracts the flag from the layer's first inbound node and bakes it into the JSON config so the runtime can pick it up. GELU gains the tanh-approximation form. A serializable gelu_approximate helper in convert_model lets Keras models reference it via activation=gelu_approximate; the converter rewrites that to a plain "gelu" activation with approximate=True in the config, which the C++ gelu_layer applies via the appropriate formula. Both features are exercised by new cases in the exhaustive test model. Co-Authored-By: Claude Opus 4.7 (1M context) --- include/fdeep/import_model.hpp | 12 ++-- include/fdeep/layers/gelu_layer.hpp | 8 ++- .../layers/multi_head_attention_layer.hpp | 32 ++++++++- include/fdeep/recurrent_ops.hpp | 9 +++ keras_export/convert_model.py | 68 +++++++++++++++++++ keras_export/generate_test_models.py | 21 ++++++ 6 files changed, 141 insertions(+), 9 deletions(-) diff --git a/include/fdeep/import_model.hpp b/include/fdeep/import_model.hpp index bd0b2093..ec981462 100644 --- a/include/fdeep/import_model.hpp +++ b/include/fdeep/import_model.hpp @@ -1329,11 +1329,11 @@ namespace internal { const get_param_f&, const nlohmann::json& data, const std::string& name) { + bool approximate = false; if (json_obj_has_member(data, "config") && json_obj_has_member(data["config"], "approximate") && !data["config"]["approximate"].is_null()) { - const bool approximate = data["config"]["approximate"]; - assertion(approximate == false, "Gelu with approximate = True is not supported."); + approximate = data["config"]["approximate"]; } - return std::make_shared(name); + return std::make_shared(name, approximate); } inline activation_layer_ptr create_softsign_layer( @@ -1473,6 +1473,9 @@ namespace internal { const std::size_t key_dim = data["config"]["key_dim"]; const std::size_t value_dim = data["config"]["value_dim"]; const bool use_bias = data["config"]["use_bias"]; + const bool use_causal_mask = json_obj_has_member(data["config"], "use_causal_mask") + ? data["config"]["use_causal_mask"].get() + : false; const auto weight_shapes = create_vector>(fplus::bind_1st_of_2( create_vector, create_size_t), get_param(name, "weight_shapes")); @@ -1485,7 +1488,8 @@ namespace internal { }, weight_shapes, weight_values); return std::make_shared(name, - num_heads, key_dim, value_dim, use_bias, weights_and_biases); + num_heads, key_dim, value_dim, use_bias, use_causal_mask, + weights_and_biases); } inline layer_ptr create_lstm_layer(const get_param_f& get_param, diff --git a/include/fdeep/layers/gelu_layer.hpp b/include/fdeep/layers/gelu_layer.hpp index 7e728729..7059340b 100644 --- a/include/fdeep/layers/gelu_layer.hpp +++ b/include/fdeep/layers/gelu_layer.hpp @@ -17,16 +17,20 @@ namespace internal { class gelu_layer : public activation_layer { public: - explicit gelu_layer(const std::string& name) + explicit gelu_layer(const std::string& name, bool approximate = false) : activation_layer(name) + , approximate_(approximate) { } protected: tensor transform_input(const tensor& in_vol) const override { - return transform_tensor(gelu_activation, in_vol); + return approximate_ + ? transform_tensor(gelu_approximate_activation, in_vol) + : transform_tensor(gelu_activation, in_vol); } + bool approximate_; }; } diff --git a/include/fdeep/layers/multi_head_attention_layer.hpp b/include/fdeep/layers/multi_head_attention_layer.hpp index 2edee0df..bb677fa7 100644 --- a/include/fdeep/layers/multi_head_attention_layer.hpp +++ b/include/fdeep/layers/multi_head_attention_layer.hpp @@ -19,11 +19,13 @@ namespace internal { public: explicit multi_head_attention_layer(const std::string& name, std::size_t num_heads, std::size_t key_dim, std::size_t value_dim, - bool use_bias, const std::vector& weights_and_biases) + bool use_bias, bool use_causal_mask, + const std::vector& weights_and_biases) : layer(name) , num_heads_(num_heads) , key_dim_(key_dim) , value_dim_(value_dim) + , use_causal_mask_(use_causal_mask) , query_dense_(create_dense_layers(weights_and_biases, use_bias, num_heads, 0, key_dim, name + "_query_dense")) , value_dense_(create_dense_layers(weights_and_biases, use_bias, num_heads, 2, value_dim, name + "_value_dense")) , key_dense_(create_dense_layers(weights_and_biases, use_bias, num_heads, 1, key_dim, name + "_key_dense")) @@ -88,12 +90,35 @@ namespace internal { // https://dmol.pub/dl/attention.html#multi-head-attention-block // https://github.com/keras-team/keras/blob/v2.14.0/keras/layers/attention/multi_head_attention.py // https://gist.github.com/sevagh/b71d253a347a9b59c026580625452fc5 - const tensor scores = dot_product_tensors(query, transpose(key), std::vector({ 2, 1 }), false); + tensor scores = dot_product_tensors(query, transpose(key), std::vector({ 2, 1 }), false); const std::size_t query_size = query.shape().depth_; - const tensor distribution = softmax(transform_tensor(fplus::multiply_with(1 / std::sqrt(query_size)), scores)); + scores = transform_tensor(fplus::multiply_with(1 / std::sqrt(query_size)), scores); + if (use_causal_mask_) { + apply_causal_mask(scores); + } + const tensor distribution = softmax(scores); return dot_product_tensors(distribution, value, std::vector({ 2, 1 }), false); } + static void apply_causal_mask(tensor& scores) + { + // Scores have shape (..., T, S) with width=T (query positions) + // and depth=S (key positions). Mask out s > t. + const auto& s = scores.shape(); + const float_type neg_inf = -std::numeric_limits::infinity(); + for (std::size_t d5 = 0; d5 < s.size_dim_5_; ++d5) { + for (std::size_t d4 = 0; d4 < s.size_dim_4_; ++d4) { + for (std::size_t y = 0; y < s.height_; ++y) { + for (std::size_t t = 0; t < s.width_; ++t) { + for (std::size_t k = t + 1; k < s.depth_; ++k) { + scores.set_ignore_rank(tensor_pos(d5, d4, y, t, k), neg_inf); + } + } + } + } + } + } + protected: tensors apply_impl(const tensors& input) const override { @@ -111,6 +136,7 @@ namespace internal { std::size_t num_heads_; std::size_t key_dim_; std::size_t value_dim_; + bool use_causal_mask_; std::vector query_dense_; std::vector value_dense_; std::vector key_dense_; diff --git a/include/fdeep/recurrent_ops.hpp b/include/fdeep/recurrent_ops.hpp index a62ce33e..7c9712bf 100644 --- a/include/fdeep/recurrent_ops.hpp +++ b/include/fdeep/recurrent_ops.hpp @@ -75,6 +75,15 @@ namespace internal { return static_cast(0.5) * x * (static_cast(1) + static_cast(std::erf(x / std::sqrt(static_cast(2))))); } + inline float_type gelu_approximate_activation(float_type x) + { + // 0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3))) + const float_type sqrt_2_over_pi = static_cast(0.7978845608028654); + const float_type c = static_cast(0.044715); + const float_type inner = sqrt_2_over_pi * (x + c * x * x * x); + return static_cast(0.5) * x * (static_cast(1) + std::tanh(inner)); + } + inline float_type softsign_activation(float_type x) { return x / (std::abs(x) + static_cast(1)); diff --git a/keras_export/convert_model.py b/keras_export/convert_model.py index 871371f1..34ab33db 100755 --- a/keras_export/convert_model.py +++ b/keras_export/convert_model.py @@ -11,11 +11,21 @@ import numpy as np import numpy.typing # pylint: disable=unused-import +import keras from keras import backend as K, Layer from keras.layers import Input, Embedding, CategoryEncoding from keras.models import Model, load_model from keras.src import Functional + +@keras.saving.register_keras_serializable(package="fdeep") +def gelu_approximate(x): # type: ignore[no-untyped-def] + """Tanh-approximation form of GELU, registered so saved models that use + it can be loaded by this converter. The conversion step rewrites it to a + plain ``gelu`` activation with an extra ``approximate=True`` flag in the + layer config, which the C++ runtime picks up.""" + return keras.activations.gelu(x, approximate=True) + __author__ = "Tobias Hermann" __copyright__ = "Copyright 2017, Tobias Hermann" __license__ = "MIT" @@ -968,6 +978,62 @@ def calculate_hash(model: Model) -> str: return hash_m.hexdigest() +_ACTIVATION_FUNCTION_REWRITES = { + 'fdeep>gelu_approximate': ('gelu', {'approximate': True}), +} + + +def rewrite_custom_activations(arch: Any) -> None: + """Walk the serialized architecture and replace registered custom + activation functions with a plain string activation plus extra config + flags, so the C++ runtime sees a known activation type.""" + if isinstance(arch, dict): + cfg = arch.get('config') + if isinstance(cfg, dict): + act = cfg.get('activation') + if isinstance(act, dict) and act.get('class_name') == 'function': + rewrite = _ACTIVATION_FUNCTION_REWRITES.get(act.get('config')) + if rewrite is not None: + new_name, extra = rewrite + cfg['activation'] = new_name + cfg.update(extra) + for v in arch.values(): + rewrite_custom_activations(v) + elif isinstance(arch, list): + for v in arch: + rewrite_custom_activations(v) + + +def inject_mha_call_kwargs(arch: Any, model: Model) -> None: + """``use_causal_mask`` is a per-call kwarg of ``MultiHeadAttention``, not a + layer config field, so ``model.to_json()`` doesn't include it. Look the + flag up from the layer's first inbound node and bake it into the JSON + config so the C++ runtime can read it.""" + if isinstance(arch, dict): + if arch.get('class_name') == 'MultiHeadAttention': + cfg = arch.get('config') or {} + layer_name = cfg.get('name') + if layer_name: + try: + layer = model.get_layer(layer_name) + except (ValueError, KeyError): + layer = None + if layer is not None and layer._inbound_nodes: + kwargs = layer._inbound_nodes[0].arguments.kwargs + if kwargs.get('use_causal_mask'): + cfg['use_causal_mask'] = True + if kwargs.get('attention_mask') is not None: + raise NotImplementedError( + f"MultiHeadAttention layer {layer_name!r} was " + "called with an explicit attention_mask, which " + "the frugally-deep runtime does not yet support.") + for v in arch.values(): + inject_mha_call_kwargs(v, model) + elif isinstance(arch, list): + for v in arch: + inject_mha_call_kwargs(v, model) + + def model_to_fdeep_json(model: Model, no_tests: bool = False) -> Mapping[str, Any]: """Convert any Keras model to the frugally-deep model format.""" @@ -983,6 +1049,8 @@ def model_to_fdeep_json(model: Model, no_tests: bool = False) -> Mapping[str, An json_output = {} print('Converting model architecture.') json_output['architecture'] = json.loads(model.to_json()) + rewrite_custom_activations(json_output['architecture']) + inject_mha_call_kwargs(json_output['architecture'], model) json_output['image_data_format'] = K.image_data_format() json_output['input_shapes'] = list(map(get_layer_input_shape_tensor_shape, get_model_input_layers(model))) json_output['output_shapes'] = list(map(keras_shape_to_fdeep_tensor_shape, as_list(model.output_shape))) diff --git a/keras_export/generate_test_models.py b/keras_export/generate_test_models.py index e142a649..fba66db1 100644 --- a/keras_export/generate_test_models.py +++ b/keras_export/generate_test_models.py @@ -5,6 +5,11 @@ import sys from typing import Tuple, List, Union +# convert_model lives next to this script; importing it pulls in the +# ``gelu_approximate`` serializable activation that the converter then +# rewrites to ``gelu`` with ``approximate=True``. +from convert_model import gelu_approximate + import numpy as np from keras import activations from keras.layers import ActivityRegularization @@ -413,6 +418,8 @@ def get_test_model_exhaustive() -> Model: outputs.append(Dense(3, use_bias=True)(inputs[14])) outputs.append(Dense(4, use_bias=False)(inputs[16])) outputs.append(Dense(4, use_bias=False, activation='tanh')(inputs[18])) + outputs.append(Dense(4, use_bias=True, activation='gelu')(inputs[18])) + outputs.append(Dense(4, use_bias=True, activation=gelu_approximate)(inputs[18])) outputs.append(Dense(4, use_bias=False)(inputs[20])) outputs.append(Reshape(((2 * 3 * 4 * 5 * 6),))(inputs[0])) @@ -541,6 +548,20 @@ def get_test_model_exhaustive() -> Model: num_heads=2, key_dim=3, value_dim=5, use_bias=True, output_shape=None, attention_axes=None)(inputs[49], inputs[50], inputs[51])) + # use_causal_mask=True: triangular attention so position i only attends to <=i. + outputs.append(MultiHeadAttention( + num_heads=1, key_dim=2, value_dim=None, + use_bias=True, output_shape=None, attention_axes=None)( + inputs[50], inputs[50], use_causal_mask=True)) + outputs.append(MultiHeadAttention( + num_heads=2, key_dim=3, value_dim=4, + use_bias=True, output_shape=None, attention_axes=None)( + inputs[50], inputs[50], inputs[50], use_causal_mask=True)) + outputs.append(MultiHeadAttention( + num_heads=3, key_dim=2, value_dim=2, + use_bias=False, output_shape=None, attention_axes=None)( + inputs[50], inputs[50], use_causal_mask=True)) + # GroupedQueryAttention: shared Q/K/V seq, separate K/V seq, with/without gate. outputs.append(GroupQueryAttention(head_dim=4, num_query_heads=6, num_key_value_heads=2)(inputs[49], inputs[49], inputs[49]))