diff --git a/gemma/activations.h b/gemma/activations.h index 724931b9..3c22d64f 100644 --- a/gemma/activations.h +++ b/gemma/activations.h @@ -197,7 +197,9 @@ struct AttentionActivations { MatStorageT vit_K_T; MatStorageT vit_V_T; - MatStorageT pre_att_rms_out; + // BF16 because this is only ever the A operand of the QKV MatMuls, which + // would otherwise decompress it to BF16 on every call; see `pre_ffw_rms_out`. + MatStorageT pre_att_rms_out; MatStorageT att_out; // attention output MatStorageT att_out_reps; // attention output for each thread. MatStorageT softmax_max; // see OnlineSoftmaxState @@ -308,7 +310,7 @@ struct AttentionActivationsPtrs { MatPtrT vit_V_T; // Output of RMSNorm before attention, size batch_size x model_dim. - MatPtrT pre_att_rms_out; + MatPtrT pre_att_rms_out; // Attention output computed from att * V, size batch_size x (q_heads * // qkv_dim). MatPtrT att_out; diff --git a/gemma/attention_test.cc b/gemma/attention_test.cc index af0ae7c6..c7c45f4c 100644 --- a/gemma/attention_test.cc +++ b/gemma/attention_test.cc @@ -43,12 +43,13 @@ HWY_BEFORE_NAMESPACE(); namespace gcpp { namespace HWY_NAMESPACE { -void FillRandom(MatPtrT& mat, uint64_t seed) { +template +void FillRandom(MatPtrT& mat, uint64_t seed) { hwy::RandomState rng(seed); for (size_t r = 0; r < mat.Rows(); ++r) { - float* row = mat.Row(r); + T* row = mat.Row(r); for (size_t c = 0; c < mat.Cols(); ++c) { - row[c] = static_cast(RandomGaussian(rng)); + row[c] = hwy::ConvertScalarTo(RandomGaussian(rng)); } } }