Skip to content

Commit e52d655

Browse files
Fireronincopybara-github
authored andcommitted
No public description
PiperOrigin-RevId: 992411459
1 parent 1c570b9 commit e52d655

14 files changed

Lines changed: 1380 additions & 98 deletions

‎CMakeLists.txt‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -62,7 +62,7 @@ if(GEMMA_ONEDNN_BRGEMM OR GEMMA_ONEDNN_MATMUL)
6262
endif()
6363
include(${CMAKE_CURRENT_LIST_DIR}/cmake/GemmaFetch.cmake)
6464

65-
gemma_fetch(highway GIT_REPOSITORY https://github.com/google/highway.git GIT_TAG 97a5dd1af1a43a8b2ccbd31e556b4139bffbdafd EXCLUDE_FROM_ALL)
65+
gemma_fetch(highway GIT_REPOSITORY https://github.com/google/highway.git GIT_TAG 353597727402dfc7b28e5b1474766e366ceafd24 EXCLUDE_FROM_ALL)
6666

6767
if(CMAKE_CXX_COMPILER_ID STREQUAL "GNU" AND
6868
CMAKE_CXX_COMPILER_VERSION VERSION_GREATER_EQUAL 15)

‎MODULE.bazel‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,7 @@ bazel_dep(name = "google_benchmark", version = "1.8.5")
1919
# Require a more recent version.
2020
git_override(
2121
module_name = "highway",
22-
commit = "97a5dd1af1a43a8b2ccbd31e556b4139bffbdafd",
22+
commit = "353597727402dfc7b28e5b1474766e366ceafd24",
2323
remote = "https://github.com/google/highway",
2424
)
2525

‎examples/hello_world/CMakeLists.txt‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -35,7 +35,7 @@ else()
3535
endfunction()
3636
endif()
3737

38-
gemma_fetch(highway GIT_REPOSITORY https://github.com/google/highway.git GIT_TAG 97a5dd1af1a43a8b2ccbd31e556b4139bffbdafd)
38+
gemma_fetch(highway GIT_REPOSITORY https://github.com/google/highway.git GIT_TAG 353597727402dfc7b28e5b1474766e366ceafd24)
3939
if(CMAKE_CXX_COMPILER_ID STREQUAL "GNU" AND
4040
CMAKE_CXX_COMPILER_VERSION VERSION_GREATER_EQUAL 15)
4141
target_compile_definitions(hwy PUBLIC HWY_DISABLED_TARGETS=HWY_AVX10_2)

‎examples/simplified_gemma/CMakeLists.txt‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -35,7 +35,7 @@ else()
3535
endfunction()
3636
endif()
3737

38-
gemma_fetch(highway GIT_REPOSITORY https://github.com/google/highway.git GIT_TAG 97a5dd1af1a43a8b2ccbd31e556b4139bffbdafd)
38+
gemma_fetch(highway GIT_REPOSITORY https://github.com/google/highway.git GIT_TAG 353597727402dfc7b28e5b1474766e366ceafd24)
3939
if(CMAKE_CXX_COMPILER_ID STREQUAL "GNU" AND
4040
CMAKE_CXX_COMPILER_VERSION VERSION_GREATER_EQUAL 15)
4141
target_compile_definitions(hwy PUBLIC HWY_DISABLED_TARGETS=HWY_AVX10_2)

‎gemma/activations.h‎

Lines changed: 34 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -131,7 +131,19 @@ struct AttentionActivations {
131131
inv_timescale_global(CreateInvTimescale(
132132
allocator, max_qkv_dim,
133133
layer_config.post_qk == PostQKType::HalfRope,
134-
config.global_rope_theta, config.partial_rotary_factor)) {
134+
config.global_rope_theta, config.partial_rotary_factor)),
135+
s_att_q(config.is_encoder_decoder ? config.decoder_num_layers
136+
: config.num_layers,
137+
max_workers),
138+
s_att_k(config.is_encoder_decoder ? config.decoder_num_layers
139+
: config.num_layers,
140+
max_workers),
141+
s_att_v(config.is_encoder_decoder ? config.decoder_num_layers
142+
: config.num_layers,
143+
max_workers),
144+
s_att_out(config.is_encoder_decoder ? config.decoder_num_layers
145+
: config.num_layers,
146+
max_workers) {
135147
// Batch size can be 0 in experimental code so do not assert.
136148
if (batch_size == 0) {
137149
static std::atomic_flag warned = ATOMIC_FLAG_INIT;
@@ -246,6 +258,13 @@ struct AttentionActivations {
246258
// Rope
247259
MatStorageT<float> inv_timescale;
248260
MatStorageT<float> inv_timescale_global;
261+
262+
// Only active when GCPP_TENSOR_STATS.
263+
TensorStats s_att_q;
264+
TensorStats s_att_k;
265+
TensorStats s_att_v;
266+
TensorStats s_att_out;
267+
249268
// Replication factor to help evenly share work over threads.
250269
static constexpr size_t kThreadReplicationFactor = 4;
251270
};
@@ -301,6 +320,10 @@ struct AttentionActivationsPtrs {
301320
int8_queries = &activations.int8_queries;
302321
float_queries = &activations.float_queries;
303322
q_scales = &activations.q_scales;
323+
s_att_q = &activations.s_att_q;
324+
s_att_k = &activations.s_att_k;
325+
s_att_v = &activations.s_att_v;
326+
s_att_out = &activations.s_att_out;
304327
}
305328

306329
void SetBatchSize(size_t batch_size) {
@@ -383,6 +406,12 @@ struct AttentionActivationsPtrs {
383406
hwy::Divisor div_heads;
384407
// Query scaling factor for attention computation.
385408
float query_scale;
409+
410+
// Only active when GCPP_TENSOR_STATS.
411+
TensorStats* s_att_q = nullptr;
412+
TensorStats* s_att_k = nullptr;
413+
TensorStats* s_att_v = nullptr;
414+
TensorStats* s_att_out = nullptr;
386415
};
387416

388417
static inline size_t MoEBatchSize(const LayerConfig& layer_config,
@@ -646,6 +675,10 @@ struct Activations {
646675
}
647676

648677
~Activations() {
678+
attention_storage.s_att_q.ReduceAndPrint("att_q");
679+
attention_storage.s_att_k.ReduceAndPrint("att_k");
680+
attention_storage.s_att_v.ReduceAndPrint("att_v");
681+
attention_storage.s_att_out.ReduceAndPrint("att_out");
649682
s_ffw_in.ReduceAndPrint("ffw_in");
650683
s_ffw_hidden.ReduceAndPrint("ffw_hidden");
651684
s_ffw_out.ReduceAndPrint("ffw_out");

‎gemma/configs.cc‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1304,6 +1304,7 @@ constexpr std::pair<const char*, AttentionImpl> kAttentionImplNameToEnum[] = {
13041304
{"flash_matrix_accumulation", AttentionImpl::kFlashMatrixAccumulation},
13051305
{"int8_matrix_accumulation", AttentionImpl::kInt8MatrixAccumulation},
13061306
{"flash_amx", AttentionImpl::kFlashAMX},
1307+
{"flash_amx_int8", AttentionImpl::kFlashAMXInt8},
13071308
};
13081309

13091310
std::string GetAttentionImplName(AttentionImpl impl) {

‎gemma/configs.h‎

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -151,6 +151,7 @@ enum class AttentionImpl {
151151
kInt8MatrixAccumulation,
152152
kFlashTransposedQsInt8,
153153
kFlashAMX,
154+
kFlashAMXInt8,
154155
kSentinel,
155156
};
156157

@@ -164,14 +165,23 @@ static inline bool IsTiledAttention(AttentionImpl impl) {
164165
impl == AttentionImpl::kFlashTransposedQsInt8 ||
165166
impl == AttentionImpl::kInt8MatrixAccumulation ||
166167
impl == AttentionImpl::kFlashMatrixAccumulation ||
167-
impl == AttentionImpl::kFlashAMX;
168+
impl == AttentionImpl::kFlashAMX ||
169+
impl == AttentionImpl::kFlashAMXInt8;
168170
}
169171

170172
static inline bool IsBF16TransposedQsAttention(AttentionImpl impl) {
171173
return impl == AttentionImpl::kFlashTransposedQsBF16 ||
172174
impl == AttentionImpl::kFlashAMX;
173175
}
174176

177+
// Int8 implementations sharing the VNNI KV tile layout (K as
178+
// [qkv_dim/4][kTileSize][4], V as [kTileSize/4][qkv_dim][4], BF16 K/V scales,
179+
// then int32 K sums) and int8 queries with per-query scales.
180+
static inline bool IsInt8VNNIAttention(AttentionImpl impl) {
181+
return impl == AttentionImpl::kFlashTransposedQsInt8 ||
182+
impl == AttentionImpl::kFlashAMXInt8;
183+
}
184+
175185
// Post attention and ffw normalization type.
176186
enum class PostNormType {
177187
None,

‎gemma/flash_attention.cc‎

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1882,6 +1882,27 @@ void DispatchTileFlashAttentionReturnExpSumsAndMaxLogitsAMX(
18821882
att_out, exp_denominator_sums, max_logits);
18831883
}
18841884

1885+
void DispatchTileFlashAttentionReturnExpSumsAndMaxLogitsAMXInt8(
1886+
hwy::Span<const MatPtr> kvs, size_t q_count,
1887+
const int8_t* HWY_RESTRICT q_base, const hwy::Span<const float> q_scales,
1888+
hwy::Span<const size_t> start_pos_per_query,
1889+
hwy::Span<const size_t> last_pos_per_query, const float att_cap,
1890+
MatPtrT<float>& att_out, float* HWY_RESTRICT exp_denominator_sums,
1891+
float* HWY_RESTRICT max_logits,
1892+
hwy::AlignedVector<uint8_t>* worker_workspace) {
1893+
TileFlashAttentionReturnExpSumsAndMaxLogitsAMXInt8(
1894+
kvs, q_count, q_base, q_scales, start_pos_per_query, last_pos_per_query,
1895+
att_cap, att_out, exp_denominator_sums, max_logits, worker_workspace);
1896+
}
1897+
1898+
bool AmxInt8Available() {
1899+
#if GEMMA_HAVE_AMX_INT8
1900+
return hwy::HaveTile64BMatMulI8();
1901+
#else
1902+
return false;
1903+
#endif
1904+
}
1905+
18851906
void DispatchTileFlashAttentionReturnExpSumsAndMaxLogitsInt16(
18861907
hwy::Span<const MatPtr> kvs, size_t q_count,
18871908
const int16_t* HWY_RESTRICT q_base, const hwy::Span<const float> q_scales,

‎gemma/flash_attention.h‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -77,6 +77,18 @@ namespace gcpp {
7777
MatPtrT<float>& att_out, float* HWY_RESTRICT exp_denominator_sums, \
7878
float* HWY_RESTRICT max_logits); \
7979
\
80+
void DispatchTileFlashAttentionReturnExpSumsAndMaxLogitsAMXInt8( \
81+
hwy::Span<const MatPtr> kvs, size_t q_count, \
82+
const int8_t* HWY_RESTRICT q_base, hwy::Span<const float> q_scales, \
83+
hwy::Span<const size_t> start_pos_per_query, \
84+
hwy::Span<const size_t> last_pos_per_query, const float att_cap, \
85+
MatPtrT<float>& att_out, float* HWY_RESTRICT exp_denominator_sums, \
86+
float* HWY_RESTRICT max_logits, \
87+
hwy::AlignedVector<uint8_t>* worker_workspace); \
88+
\
89+
/* True if the AMX-INT8 kernel is compiled in and usable on this CPU. */ \
90+
bool AmxInt8Available(); \
91+
\
8092
void DispatchTileFlashAttentionReturnExpSumsAndMaxLogitsInt16( \
8193
hwy::Span<const MatPtr> kvs, size_t q_count, \
8294
const int16_t* HWY_RESTRICT q_base, hwy::Span<const float> q_scales, \

0 commit comments

Comments
 (0)