///|
fn materialized_model_shape(
dimensions : Array[Int],
) -> @shape.Shape raise @tensor.TensorError {
@shape.Shape::new(dimensions) catch {
error => raise @tensor.TensorError::new(error.to_string())
}
}
///|
fn WebNNGraphBuilder::materialize_attention_mask(
self : WebNNGraphBuilder,
mask : @model.AttentionMask?,
) -> WebNNTensor? raise @tensor.TensorError {
match mask {
Some(data) => Some(self.constant(data.shape(), data.values()))
None => None
}
}
///|
fn WebNNGraphBuilder::materialize_self_attention_with_mask(
self : WebNNGraphBuilder,
config : @model.SelfAttentionConfig,
parameters : @model.SelfAttentionParameters,
additive_mask : WebNNTensor?,
) -> @model.SelfAttention[WebNNTensor] raise @tensor.TensorError {
if !parameters.matches(config) {
raise @tensor.TensorError::new(
"attention parameters do not match attention config",
)
}
let width = config.width()
let weight_shape = materialized_model_shape([width, width])
let bias_shape = materialized_model_shape([width])
{
query: {
weight: self.constant(weight_shape, parameters.query_weight()),
bias: self.constant(bias_shape, parameters.query_bias()),
},
key: {
weight: self.constant(weight_shape, parameters.key_weight()),
bias: self.constant(bias_shape, parameters.key_bias()),
},
value: {
weight: self.constant(weight_shape, parameters.value_weight()),
bias: self.constant(bias_shape, parameters.value_bias()),
},
output: {
weight: self.constant(weight_shape, parameters.output_weight()),
bias: self.constant(bias_shape, parameters.output_bias()),
},
score_scale: self.constant(materialized_model_shape([1]), [
config.score_scale(),
]),
additive_mask,
heads: config.heads(),
}
}
///|
pub fn WebNNGraphBuilder::materialize_self_attention(
self : WebNNGraphBuilder,
config : @model.SelfAttentionConfig,
parameters : @model.SelfAttentionParameters,
mask : @model.AttentionMask?,
) -> @model.SelfAttention[WebNNTensor] raise @tensor.TensorError {
self.materialize_self_attention_with_mask(
config,
parameters,
self.materialize_attention_mask(mask),
)
}
///|
pub fn WebNNGraphBuilder::materialize_feed_forward(
self : WebNNGraphBuilder,
config : @model.TransformerEncoderConfig,
parameters : @model.FeedForwardParameters,
) -> @model.FeedForward[WebNNTensor] raise @tensor.TensorError {
if !parameters.matches(config) {
raise @tensor.TensorError::new(
"feed-forward parameters do not match encoder config",
)
}
let width = config.width()
let hidden_size = config.hidden_size()
{
input_projection: {
weight: self.constant(
materialized_model_shape([width, hidden_size]),
parameters.input_weight(),
),
bias: self.constant(
materialized_model_shape([hidden_size]),
parameters.input_bias(),
),
},
output_projection: {
weight: self.constant(
materialized_model_shape([hidden_size, width]),
parameters.output_weight(),
),
bias: self.constant(
materialized_model_shape([width]),
parameters.output_bias(),
),
},
}
}
///|
fn WebNNGraphBuilder::materialize_transformer_encoder_with_mask(
self : WebNNGraphBuilder,
config : @model.TransformerEncoderConfig,
parameters : @model.TransformerEncoderParameters,
additive_mask : WebNNTensor?,
) -> @model.TransformerEncoderBlock[WebNNTensor] raise @tensor.TensorError {
let width_shape = materialized_model_shape([config.width()])
{
attention_normalization: {
scale: self.constant(
width_shape,
parameters.attention_normalization_scale(),
),
bias: self.constant(
width_shape,
parameters.attention_normalization_bias(),
),
axes: [config.normalization_axis()],
epsilon: config.epsilon(),
},
attention: self.materialize_self_attention_with_mask(
config.attention(),
parameters.attention(),
additive_mask,
),
feed_forward_normalization: {
scale: self.constant(
width_shape,
parameters.feed_forward_normalization_scale(),
),
bias: self.constant(
width_shape,
parameters.feed_forward_normalization_bias(),
),
axes: [config.normalization_axis()],
epsilon: config.epsilon(),
},
feed_forward: self.materialize_feed_forward(
config,
parameters.feed_forward(),
),
}
}
///|
pub fn WebNNGraphBuilder::materialize_transformer_encoder(
self : WebNNGraphBuilder,
config : @model.TransformerEncoderConfig,
parameters : @model.TransformerEncoderParameters,
mask : @model.AttentionMask?,
) -> @model.TransformerEncoderBlock[WebNNTensor] raise @tensor.TensorError {
self.materialize_transformer_encoder_with_mask(
config,
parameters,
self.materialize_attention_mask(mask),
)
}
///|
pub fn WebNNGraphBuilder::materialize_transformer_encoder_stack(
self : WebNNGraphBuilder,
config : @model.TransformerEncoderStackConfig,
parameters : @model.TransformerEncoderStackParameters,
mask : @model.AttentionMask?,
) -> @model.TransformerEncoderStack[WebNNTensor] raise @tensor.TensorError {
let encoder_config = config.encoder()
let additive_mask = self.materialize_attention_mask(mask)
let layer_parameters = parameters.layers()
if layer_parameters.length() != config.layers() {
raise @tensor.TensorError::new(
"encoder stack parameters do not match stack config",
)
}
let layers : Array[@model.TransformerEncoderBlock[WebNNTensor]] = []
for parameters in layer_parameters {
layers.push(
self.materialize_transformer_encoder_with_mask(
encoder_config, parameters, additive_mask,
),
)
}
@model.TransformerEncoderStack::new(layers)
}
///|
fn WebNNGraphBuilder::materialize_bert_encoder_with_mask(
self : WebNNGraphBuilder,
config : @model.TransformerEncoderConfig,
parameters : @model.TransformerEncoderParameters,
additive_mask : WebNNTensor?,
) -> @model.BertEncoderBlock[WebNNTensor] raise @tensor.TensorError {
if !parameters.matches(config) {
raise @tensor.TensorError::new(
"BERT encoder parameters do not match encoder config",
)
}
let width_shape = @shape.Shape::new([config.width()]) catch {
error => raise @tensor.TensorError::new(error.to_string())
}
{
attention: self.materialize_self_attention_with_mask(
config.attention(),
parameters.attention(),
additive_mask,
),
attention_output_normalization: {
scale: self.constant(
width_shape,
parameters.attention_normalization_scale(),
),
bias: self.constant(
width_shape,
parameters.attention_normalization_bias(),
),
axes: [config.normalization_axis()],
epsilon: config.epsilon(),
},
feed_forward: self.materialize_feed_forward(
config,
parameters.feed_forward(),
),
output_normalization: {
scale: self.constant(
width_shape,
parameters.feed_forward_normalization_scale(),
),
bias: self.constant(
width_shape,
parameters.feed_forward_normalization_bias(),
),
axes: [config.normalization_axis()],
epsilon: config.epsilon(),
},
}
}
///|
pub fn WebNNGraphBuilder::materialize_bert_encoder_stack(
self : WebNNGraphBuilder,
config : @model.BertEncoderConfig,
parameters : @model.BertEncoderParameters,
mask : @model.AttentionMask?,
) -> @model.BertEncoderStack[WebNNTensor] raise @tensor.TensorError {
let encoder_config = config.encoder()
let additive_mask = self.materialize_attention_mask(mask)
let layer_parameters = parameters.layers()
if layer_parameters.length() != config.layers() {
raise @tensor.TensorError::new(
"BERT encoder parameters do not match stack config",
)
}
let layers : Array[@model.BertEncoderBlock[WebNNTensor]] = []
for parameters in layer_parameters {
layers.push(
self.materialize_bert_encoder_with_mask(
encoder_config, parameters, additive_mask,
),
)
}
@model.BertEncoderStack::new(layers)
}