| 219 | } |
| 220 | relative_position = std::llabs(relative_position); |
| 221 | } else { |
| 222 | relative_position = -std::min<int64_t>(relative_position, 0); |
| 223 | } |
| 224 | const int64_t max_exact = buckets / 2; |
| 225 | int64_t bucket = relative_position; |
| 226 | if (relative_position >= max_exact) { |
| 227 | const double log_ratio = std::log(static_cast<double>(relative_position) / static_cast<double>(max_exact)) / |
| 228 | std::log(static_cast<double>(max_distance) / static_cast<double>(max_exact)); |
| 229 | bucket = max_exact + static_cast<int64_t>(log_ratio * static_cast<double>(buckets - max_exact)); |
| 230 | bucket = std::min(bucket, buckets - 1); |
| 231 | } |
| 232 | relative_buckets += bucket; |
| 233 | out[static_cast<size_t>(q * key_length + k)] = static_cast<int32_t>(relative_buckets); |
| 234 | } |
| 235 | } |
| 236 | return out; |
| 237 | } |
| 238 | |
| 239 | T5BaseEncoderModule::T5BaseEncoderModule(T5BaseEncoderConfig config) : config_(config) { |
| 240 | validate_config(config_); |
| 241 | } |
| 242 | |
| 243 | const T5BaseEncoderConfig & T5BaseEncoderModule::config() const noexcept { |
| 244 | return config_; |
| 245 | } |
| 246 | |
| 247 | core::TensorValue T5BaseEncoderModule::build( |
| 248 | core::ModuleBuildContext & ctx, |
| 249 | const core::TensorValue & input_ids, |
no test coverage detected