use objc2::rc::Retained;
use objc2::runtime::ProtocolObject;
use objc2_foundation::{NSArray, NSRange, NSString};
use objc2_metal::{
MTLCommonCounterSetTimestamp, MTLComputePassDescriptor, MTLCounterResultTimestamp,
MTLCounterSampleBuffer, MTLCounterSampleBufferDescriptor, MTLCounterSamplingPoint,
MTLCounterSet, MTLDevice, MTLRenderPassDescriptor, MTLStorageMode,
};
pub(super) use concinnity_core::render::render_graph::{PASS_COUNT, PASS_NAMES, PassId};
pub(super) const FRAMES_IN_FLIGHT: usize = 3;
const NO_SAMPLE: usize = usize::MAX;
const _: () = assert!(PASS_COUNT <= u64::BITS as usize);
const SAMPLE_COUNT: usize = PASS_COUNT * 2;
pub(super) struct SendableSampleBuf(
pub(super) Retained<ProtocolObject<dyn MTLCounterSampleBuffer>>,
);
unsafe impl Send for SendableSampleBuf {}
pub(super) struct PassTimingResources {
buffers: [Retained<ProtocolObject<dyn MTLCounterSampleBuffer>>; FRAMES_IN_FLIGHT],
frame_slot: usize,
attached: std::sync::atomic::AtomicU64,
}
impl PassTimingResources {
pub(super) fn new(device: &ProtocolObject<dyn MTLDevice>) -> Option<Self> {
if !device.supportsCounterSampling(MTLCounterSamplingPoint::AtStageBoundary) {
return None;
}
let sets: Retained<NSArray<ProtocolObject<dyn MTLCounterSet>>> = device.counterSets()?;
let target_name = unsafe { MTLCommonCounterSetTimestamp };
let mut timestamp_set: Option<Retained<ProtocolObject<dyn MTLCounterSet>>> = None;
for set in sets.iter() {
let name: Retained<NSString> = set.name();
if &*name == target_name {
timestamp_set = Some(set);
break;
}
}
let timestamp_set = timestamp_set?;
let make_buf = || -> Option<Retained<ProtocolObject<dyn MTLCounterSampleBuffer>>> {
let desc = MTLCounterSampleBufferDescriptor::new();
desc.setCounterSet(Some(×tamp_set));
desc.setStorageMode(MTLStorageMode::Shared);
unsafe { desc.setSampleCount(SAMPLE_COUNT) };
device
.newCounterSampleBufferWithDescriptor_error(&desc)
.ok()
};
let b0 = make_buf()?;
let b1 = make_buf()?;
let b2 = make_buf()?;
Some(Self {
buffers: [b0, b1, b2],
frame_slot: 0,
attached: std::sync::atomic::AtomicU64::new(0),
})
}
fn active(&self) -> &ProtocolObject<dyn MTLCounterSampleBuffer> {
&self.buffers[self.frame_slot]
}
pub(super) fn begin_frame(&mut self) -> usize {
let slot = self.frame_slot;
self.frame_slot = (self.frame_slot + 1) % FRAMES_IN_FLIGHT;
self.attached.store(0, std::sync::atomic::Ordering::Relaxed);
slot
}
pub(super) fn attached_mask(&self) -> u64 {
self.attached.load(std::sync::atomic::Ordering::Relaxed)
}
fn mark_attached(&self, pass: PassId) {
self.attached.fetch_or(
1u64 << (pass as usize),
std::sync::atomic::Ordering::Relaxed,
);
}
pub(super) fn buffer_for(
&self,
slot: usize,
) -> Retained<ProtocolObject<dyn MTLCounterSampleBuffer>> {
self.buffers[slot].clone()
}
pub(super) fn attach_render(&self, desc: &MTLRenderPassDescriptor, pass: PassId) {
self.mark_attached(pass);
let (start, end) = slot_pair(pass);
unsafe {
let arr = desc.sampleBufferAttachments();
let entry = arr.objectAtIndexedSubscript(0);
entry.setSampleBuffer(Some(self.active()));
entry.setStartOfVertexSampleIndex(NO_SAMPLE);
entry.setEndOfVertexSampleIndex(NO_SAMPLE);
entry.setStartOfFragmentSampleIndex(start);
entry.setEndOfFragmentSampleIndex(end);
}
}
pub(super) fn attach_render_first(&self, desc: &MTLRenderPassDescriptor, pass: PassId) {
self.mark_attached(pass);
let (start, _) = slot_pair(pass);
unsafe {
let arr = desc.sampleBufferAttachments();
let entry = arr.objectAtIndexedSubscript(0);
entry.setSampleBuffer(Some(self.active()));
entry.setStartOfVertexSampleIndex(NO_SAMPLE);
entry.setEndOfVertexSampleIndex(NO_SAMPLE);
entry.setStartOfFragmentSampleIndex(start);
entry.setEndOfFragmentSampleIndex(NO_SAMPLE);
}
}
pub(super) fn attach_render_last(&self, desc: &MTLRenderPassDescriptor, pass: PassId) {
let (_, end) = slot_pair(pass);
unsafe {
let arr = desc.sampleBufferAttachments();
let entry = arr.objectAtIndexedSubscript(0);
entry.setSampleBuffer(Some(self.active()));
entry.setStartOfVertexSampleIndex(NO_SAMPLE);
entry.setEndOfVertexSampleIndex(NO_SAMPLE);
entry.setStartOfFragmentSampleIndex(NO_SAMPLE);
entry.setEndOfFragmentSampleIndex(end);
}
}
pub(super) fn attach_compute(&self, desc: &MTLComputePassDescriptor, pass: PassId) {
self.mark_attached(pass);
let (start, end) = slot_pair(pass);
self.sample_compute(desc, start, end);
}
pub(super) fn attach_compute_first(&self, desc: &MTLComputePassDescriptor, pass: PassId) {
self.mark_attached(pass);
let (start, _) = slot_pair(pass);
self.sample_compute(desc, start, NO_SAMPLE);
}
pub(super) fn attach_compute_last(&self, desc: &MTLComputePassDescriptor, pass: PassId) {
let (_, end) = slot_pair(pass);
self.sample_compute(desc, NO_SAMPLE, end);
}
fn sample_compute(&self, desc: &MTLComputePassDescriptor, start: usize, end: usize) {
unsafe {
let arr = desc.sampleBufferAttachments();
let entry = arr.objectAtIndexedSubscript(0);
entry.setSampleBuffer(Some(self.active()));
entry.setStartOfEncoderSampleIndex(start);
entry.setEndOfEncoderSampleIndex(end);
}
}
pub(super) fn attach_render_timer(&self, desc: &MTLRenderPassDescriptor, timer: PassTimer) {
match timer {
PassTimer::None => {}
PassTimer::Whole(id) => self.attach_render(desc, id),
PassTimer::First(id) => self.attach_render_first(desc, id),
PassTimer::Last(id) => self.attach_render_last(desc, id),
}
}
pub(super) fn attach_compute_timer(&self, desc: &MTLComputePassDescriptor, timer: PassTimer) {
match timer {
PassTimer::None => {}
PassTimer::Whole(id) => self.attach_compute(desc, id),
PassTimer::First(id) => self.attach_compute_first(desc, id),
PassTimer::Last(id) => self.attach_compute_last(desc, id),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum PassTimer {
None,
Whole(PassId),
First(PassId),
Last(PassId),
}
impl PassTimer {
pub(crate) fn span(pass: PassId, index: usize, count: usize) -> PassTimer {
match (index == 0, index + 1 == count) {
(true, true) => PassTimer::Whole(pass),
(true, false) => PassTimer::First(pass),
(false, true) => PassTimer::Last(pass),
(false, false) => PassTimer::None,
}
}
}
fn slot_pair(pass: PassId) -> (usize, usize) {
debug_assert!(
(pass as usize) < PASS_COUNT,
"PassId {pass:?} (index {}) >= PASS_COUNT {PASS_COUNT}: register it in PASS_NAMES \
and bump PASS_COUNT",
pass as usize,
);
let base = pass as usize * 2;
(base, base + 1)
}
fn read_samples(
buffer: &ProtocolObject<dyn MTLCounterSampleBuffer>,
) -> Option<[u64; SAMPLE_COUNT]> {
let range = NSRange::new(0, SAMPLE_COUNT);
let data = unsafe { buffer.resolveCounterRange(range) }?;
let needed = std::mem::size_of::<MTLCounterResultTimestamp>() * SAMPLE_COUNT;
if data.len() < needed {
return None;
}
let timestamps: &[MTLCounterResultTimestamp] = unsafe {
std::slice::from_raw_parts(
data.as_bytes_unchecked().as_ptr() as *const MTLCounterResultTimestamp,
SAMPLE_COUNT,
)
};
let mut out = [0u64; SAMPLE_COUNT];
for (slot, sample) in out.iter_mut().zip(timestamps) {
*slot = sample.timestamp;
}
Some(out)
}
fn durations_us(samples: &[u64; SAMPLE_COUNT], attached: u64) -> [u32; PASS_COUNT] {
let mut out = [0u32; PASS_COUNT];
for (i, slot) in out.iter_mut().enumerate() {
if !ran_this_frame(samples, attached, i) {
continue;
}
let (start, end) = (samples[i * 2], samples[i * 2 + 1]);
*slot = ((end - start) / 1000).min(u32::MAX as u64) as u32;
}
out
}
fn ran_this_frame(samples: &[u64; SAMPLE_COUNT], attached: u64, i: usize) -> bool {
if attached & (1u64 << i) == 0 {
return false;
}
let (start, end) = (samples[i * 2], samples[i * 2 + 1]);
start != 0 && end != 0 && end > start
}
fn span_us(samples: &[u64; SAMPLE_COUNT], attached: u64) -> Option<u32> {
let mut min_start = u64::MAX;
let mut max_end = 0u64;
for i in 0..PASS_COUNT {
if !ran_this_frame(samples, attached, i) {
continue;
}
min_start = min_start.min(samples[i * 2]);
max_end = max_end.max(samples[i * 2 + 1]);
}
if max_end <= min_start {
return None;
}
Some(((max_end - min_start) / 1000).min(u32::MAX as u64) as u32)
}
pub(super) fn resolve(
buffer: &ProtocolObject<dyn MTLCounterSampleBuffer>,
attached: u64,
) -> [u32; PASS_COUNT] {
match read_samples(buffer) {
Some(samples) => durations_us(&samples, attached),
None => [0; PASS_COUNT],
}
}
pub(super) fn frame_span_us(
buffer: &ProtocolObject<dyn MTLCounterSampleBuffer>,
attached: u64,
) -> Option<u32> {
span_us(&read_samples(buffer)?, attached)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_span_opens_on_its_first_encoder_and_closes_on_its_last() {
let p = PassId::PlanarReflection;
assert_eq!(PassTimer::span(p, 0, 1), PassTimer::Whole(p));
let four: Vec<_> = (0..4).map(|i| PassTimer::span(p, i, 4)).collect();
assert_eq!(
four,
[
PassTimer::First(p),
PassTimer::None,
PassTimer::None,
PassTimer::Last(p)
]
);
}
fn write(samples: &mut [u64; SAMPLE_COUNT], pass: PassId, start_ns: u64, end_ns: u64) {
let (s, e) = slot_pair(pass);
samples[s] = start_ns;
samples[e] = end_ns;
}
const CLOCK: u64 = 1_700_000_000_000;
fn serial_frame() -> [u64; SAMPLE_COUNT] {
let mut s = [0u64; SAMPLE_COUNT];
let mut at = |pass, start: u64, end: u64| write(&mut s, pass, CLOCK + start, CLOCK + end);
at(PassId::Cull, 0, 30_000);
at(PassId::GBufferPrepass, 1_400_000, 1_800_000);
at(PassId::SsaoKernel, 1_800_000, 3_800_000);
at(PassId::SsaoBlur, 3_800_000, 4_200_000);
at(PassId::Shadow, 4_200_000, 5_600_000);
at(PassId::Main, 6_700_000, 11_600_000);
at(PassId::Composite, 11_600_000, 13_800_000);
s
}
fn serial_mask() -> u64 {
[
PassId::Cull,
PassId::GBufferPrepass,
PassId::SsaoKernel,
PassId::SsaoBlur,
PassId::Shadow,
PassId::Main,
PassId::Composite,
]
.iter()
.fold(0u64, |m, p| m | 1u64 << (*p as usize))
}
const GRAPHICS: [PassId; 6] = [
PassId::GBufferPrepass,
PassId::SsaoKernel,
PassId::SsaoBlur,
PassId::Shadow,
PassId::Main,
PassId::Composite,
];
fn sum_us(durations: &[u32; PASS_COUNT], passes: &[PassId]) -> u64 {
passes.iter().map(|p| durations[*p as usize] as u64).sum()
}
#[test]
fn a_pair_resolves_to_its_own_delta() {
let d = durations_us(&serial_frame(), serial_mask());
assert_eq!(d[PassId::Main as usize], 4_900);
assert_eq!(d[PassId::Shadow as usize], 1_400);
assert_eq!(d[PassId::Cull as usize], 30);
}
#[test]
fn an_unwritten_pair_reports_zero() {
let mut s = serial_frame();
assert_eq!(durations_us(&s, serial_mask())[PassId::Fog as usize], 0);
let (start, _) = slot_pair(PassId::Fog);
s[start] = CLOCK + 9_000_000;
let mask = serial_mask() | 1u64 << (PassId::Fog as usize);
assert_eq!(durations_us(&s, mask)[PassId::Fog as usize], 0);
}
#[test]
fn the_graphics_queue_sum_is_bounded_by_the_span() {
let samples = serial_frame();
let d = durations_us(&samples, serial_mask());
let span = span_us(&samples, serial_mask()).expect("a written frame has a span") as u64;
let sum = sum_us(&d, &GRAPHICS);
assert!(
sum <= span,
"graphics passes sum to {sum} us over a {span} us span"
);
assert!(
sum * 10 >= span * 8,
"sum {sum} us is not close to {span} us"
);
}
#[test]
fn the_span_covers_both_queues() {
let samples = serial_frame();
assert_eq!(span_us(&samples, serial_mask()), Some(13_800));
}
#[test]
fn a_frame_with_no_sample_has_no_span() {
assert_eq!(span_us(&[0u64; SAMPLE_COUNT], u64::MAX), None);
assert_eq!(
durations_us(&[0u64; SAMPLE_COUNT], u64::MAX),
[0; PASS_COUNT]
);
}
#[test]
fn a_pass_that_stopped_running_is_left_out_of_both_readings() {
let mut s = serial_frame();
write(
&mut s,
PassId::Ssgi,
CLOCK - 9_000_000_000,
CLOCK - 8_998_000_000,
);
let mask = serial_mask();
assert_eq!(durations_us(&s, mask)[PassId::Ssgi as usize], 0);
assert_eq!(
span_us(&s, mask),
Some(13_800),
"stale slot stretched the span"
);
let live = mask | 1u64 << (PassId::Ssgi as usize);
assert_eq!(durations_us(&s, live)[PassId::Ssgi as usize], 2_000);
assert!(span_us(&s, live).unwrap() > 9_000_000);
}
#[test]
fn every_pass_owns_a_distinct_pair_inside_the_buffer() {
let mut seen = std::collections::HashSet::new();
for pass in PassId::ALL {
let (s, e) = slot_pair(pass);
assert!(seen.insert(s), "duplicate start slot for {pass:?}");
assert!(seen.insert(e), "duplicate end slot for {pass:?}");
assert!(e < SAMPLE_COUNT, "{pass:?} addresses past the buffer");
}
}
}