| 59 | ptx_compilation_opts_(ptx_compilation_opts) {} |
| 60 | |
| 61 | port::StatusOr<DeviceMemory<uint8>> RedzoneAllocator::AllocateBytes( |
| 62 | int64 byte_size) { |
| 63 | CHECK_GE(byte_size, 0) << "byte_size must be positive."; |
| 64 | if (byte_size > GetMemoryLimitInBytes()) { |
| 65 | return port::Status( |
| 66 | port::error::RESOURCE_EXHAUSTED, |
| 67 | absl::StrFormat( |
| 68 | "Allocating %d bytes exceeds the memory limit of %d bytes.", |
| 69 | byte_size, GetMemoryLimitInBytes())); |
| 70 | } |
| 71 | |
| 72 | int64 rhs_slop = RoundUpToNearest(byte_size, kRhsRedzoneAlign) - byte_size; |
| 73 | TF_ASSIGN_OR_RETURN( |
| 74 | OwningDeviceMemory allocated_buffer, |
| 75 | memory_allocator_->Allocate(device_ordinal_, |
| 76 | byte_size + 2 * redzone_size_ + rhs_slop, |
| 77 | /*retry_on_failure=*/false)); |
| 78 | allocated_bytes_excluding_redzones_ += byte_size; |
| 79 | |
| 80 | static_assert(sizeof(uint8) == 1, "Unexpected size"); |
| 81 | DeviceMemory<uint8> allocated_buffer_memory(*allocated_buffer); |
| 82 | |
| 83 | DeviceMemory<uint8> lhs_redzone = stream_->parent()->GetSubBuffer( |
| 84 | &allocated_buffer_memory, 0, redzone_size_); |
| 85 | |
| 86 | DeviceMemory<uint8> data_chunk = stream_->parent()->GetSubBuffer( |
| 87 | &allocated_buffer_memory, redzone_size_, byte_size); |
| 88 | |
| 89 | // Split up the RHS redzone into two pieces: |
| 90 | // - 0 to kRhsRedzoneAlign bytes adjacent to the user buffer, followed by |
| 91 | // - redzone_size_ bytes. |
| 92 | // We do this because Stream::ThenMemset32 requires the buffer address and |
| 93 | // size to be aligned to 4 bytes. |
| 94 | DeviceMemory<uint8> rhs_redzone_slop = stream_->parent()->GetSubBuffer( |
| 95 | &allocated_buffer_memory, redzone_size_ + byte_size, rhs_slop); |
| 96 | |
| 97 | DeviceMemory<uint8> rhs_redzone_nonslop = stream_->parent()->GetSubBuffer( |
| 98 | &allocated_buffer_memory, redzone_size_ + byte_size + rhs_slop, |
| 99 | redzone_size_); |
| 100 | |
| 101 | uint8 pattern_arr[] = {redzone_pattern_, redzone_pattern_, redzone_pattern_, |
| 102 | redzone_pattern_}; |
| 103 | uint32 pattern32; |
| 104 | std::memcpy(&pattern32, pattern_arr, sizeof(pattern32)); |
| 105 | stream_->ThenMemset32(&lhs_redzone, pattern32, redzone_size_); |
| 106 | if (rhs_slop != 0) { |
| 107 | stream_->ThenMemcpy(&rhs_redzone_slop, &pattern32, rhs_slop); |
| 108 | } |
| 109 | stream_->ThenMemset32(&rhs_redzone_nonslop, pattern32, redzone_size_); |
| 110 | |
| 111 | allocated_buffers_.emplace_back(std::move(allocated_buffer), byte_size); |
| 112 | return data_chunk; |
| 113 | } |
| 114 | |
| 115 | // PTX blob for the function which checks that every byte in |
| 116 | // input_buffer (length is buffer_length) is equal to redzone_pattern. |