elif mode in ("mean", "mean_sqrt_len_tokens"): if mean_sum is None: mask = attention_mask.unsqueeze(-1).expand_as(token_embeddings).to(token_embeddings.dtype) mean_sum = (token_embeddings * mask).sum(dim=1) # ... mean_mask is the count of real tokens (or a supplied per-token weight sum), clamped away from zero if mode == "mean": output_vectors.append(mean_sum / mean_mask)