Skip to main content

gwr_models/memory/memory_access_gen/
strided.rs

1// Copyright (c) 2025 Graphcore Ltd. All rights reserved.
2
3use std::rc::Rc;
4
5use gwr_engine::types::AccessType;
6use gwr_track::entity::Entity;
7
8use crate::memory::memory_access::MemoryAccess;
9use crate::memory::memory_map::MemoryMap;
10
11/// A Strided address access generator.
12///
13/// Will emit memory accesses in the range [base, end)
14pub struct Strided {
15    entity: Rc<Entity>,
16    memory_map: Rc<MemoryMap>,
17
18    // Configuration
19    src_addr: u64,
20    base_addr: u64,
21    end_addr: u64,
22    stride_bytes: u64,
23    overhead_size_bytes: usize,
24    access_size_bytes: usize,
25    num_to_send: usize,
26
27    // State
28    next_addr: u64,
29    num_sent: usize,
30}
31
32impl Strided {
33    #[expect(clippy::too_many_arguments)]
34    #[must_use]
35    pub fn new(
36        parent: &Rc<Entity>,
37        name: &str,
38        memory_map: &Rc<MemoryMap>,
39        src_addr: u64,
40        base_addr: u64,
41        end_addr: u64,
42        stride_bytes: u64,
43        overhead_size_bytes: usize,
44        access_size_bytes: usize,
45        num_to_send: usize,
46    ) -> Self {
47        Self {
48            entity: Rc::new(Entity::new(parent, name)),
49            memory_map: memory_map.clone(),
50            src_addr,
51            base_addr,
52            end_addr,
53            stride_bytes,
54            overhead_size_bytes,
55            access_size_bytes,
56            num_to_send,
57            next_addr: base_addr,
58            num_sent: 0,
59        }
60    }
61}
62
63impl Iterator for Strided {
64    type Item = MemoryAccess;
65    fn next(&mut self) -> Option<Self::Item> {
66        if self.num_sent < self.num_to_send {
67            self.num_sent += 1;
68
69            let (dst_device, _) = self.memory_map.lookup(self.next_addr)?;
70            let (src_device, _) = self.memory_map.lookup(self.src_addr)?;
71
72            let access = MemoryAccess::new(
73                &self.entity,
74                AccessType::ReadRequest,
75                self.access_size_bytes,
76                self.next_addr,
77                self.src_addr,
78                dst_device,
79                src_device,
80                self.overhead_size_bytes,
81            );
82
83            self.next_addr += self.stride_bytes;
84            if self.next_addr >= self.end_addr {
85                self.next_addr = self.base_addr;
86            }
87
88            Some(access)
89        } else {
90            None
91        }
92    }
93}