| 592 | } |
| 593 | |
| 594 | ConvCache::CircularView ConvCache::get_window(size_t layer) const { |
| 595 | CircularView view; |
| 596 | if (layer >= num_layers) { |
| 597 | view.ptr1 = nullptr; |
| 598 | view.len1 = 0; |
| 599 | view.ptr2 = nullptr; |
| 600 | view.len2 = 0; |
| 601 | view.total_len = 0; |
| 602 | return view; |
| 603 | } |
| 604 | |
| 605 | const auto& state = layer_states[layer]; |
| 606 | if (state.count == 0) { |
| 607 | view.ptr1 = nullptr; |
| 608 | view.len1 = 0; |
| 609 | view.ptr2 = nullptr; |
| 610 | view.len2 = 0; |
| 611 | view.total_len = 0; |
| 612 | return view; |
| 613 | } |
| 614 | |
| 615 | size_t stride = hidden_size * element_size; |
| 616 | |
| 617 | if (state.count < window_size) { |
| 618 | view.ptr1 = state.data.data(); |
| 619 | view.len1 = state.count; |
| 620 | view.ptr2 = nullptr; |
| 621 | view.len2 = 0; |
| 622 | view.total_len = state.count; |
| 623 | return view; |
| 624 | } |
| 625 | |
| 626 | view.ptr1 = state.data.data(); |
| 627 | view.len1 = state.head; |
| 628 | view.ptr2 = state.data.data() + state.head * stride; |
| 629 | view.len2 = window_size - state.head; |
| 630 | view.total_len = window_size; |
| 631 | return view; |
| 632 | } |
| 633 | |
| 634 | void ConvCache::update(CactusGraph* gb, size_t layer, const size_t bx_node) { |
| 635 | if (layer >= num_layers || !bx_node || window_size == 0 || hidden_size == 0) { |
no test coverage detected