#![deny(unsafe_op_in_unsafe_fn)]
use concinnity_core::render::depth::CAMERA_DEPTH;
use concinnity_core::render::error::{RenderError, RenderResult};
use concinnity_core::render::history_reset::UpscalerResetLatch;
use objc2::rc::Retained;
use objc2::runtime::ProtocolObject;
use objc2_metal::{MTLDevice as _, MTLPixelFormat, MTLTexture, MTLTextureUsage};
use objc2_metal_fx::{MTLFXTemporalScaler, MTLFXTemporalScalerBase, MTLFXTemporalScalerDescriptor};
use crate::metal::context::MtlContext;
use crate::metal::descriptors::TextureDesc;
use crate::metal::error::allocation_failed;
use crate::metal::texture::{REACTIVE_MASK_FORMAT, REACTIVE_MASK_USAGE};
pub(crate) struct UpscaleState {
pub scaler: Option<MetalFXUpscaler>,
pub scale: f32,
pub jitter: UpscaleJitter,
pub reset: UpscalerResetLatch,
}
pub(crate) struct MetalFXUpscaler {
pub(crate) scaler: Retained<ProtocolObject<dyn MTLFXTemporalScaler>>,
pub(crate) output: Retained<ProtocolObject<dyn MTLTexture>>,
pub(crate) input_width: u32,
pub(crate) input_height: u32,
pub(crate) output_width: u32,
pub(crate) output_height: u32,
pub(crate) reactive: bool,
}
pub(crate) fn temporal_scaler_supported(
device: &ProtocolObject<dyn objc2_metal::MTLDevice>,
) -> bool {
unsafe { MTLFXTemporalScalerDescriptor::supportsDevice(device) }
}
fn scaler_input_size(output: (u32, u32), scale: f32, ratio_range: (f32, f32)) -> (u32, u32) {
let (min_ratio, max_ratio) = ratio_range;
let requested_ratio = if scale > 0.0 { 1.0 / scale } else { 1.0 };
let lo = min_ratio.max(1.0);
let clamped_ratio = requested_ratio.clamp(lo, max_ratio.max(lo));
let scale = 1.0 / clamped_ratio;
(
((output.0 as f32) * scale).max(1.0) as u32,
((output.1 as f32) * scale).max(1.0) as u32,
)
}
impl MetalFXUpscaler {
pub(crate) fn new(
device: &ProtocolObject<dyn objc2_metal::MTLDevice>,
output_width: u32,
output_height: u32,
scale: f32,
) -> RenderResult<Self> {
let min_ratio = unsafe {
MTLFXTemporalScalerDescriptor::supportedInputContentMinScaleForDevice(device)
};
let max_ratio = unsafe {
MTLFXTemporalScalerDescriptor::supportedInputContentMaxScaleForDevice(device)
};
let (input_width, input_height) =
scaler_input_size((output_width, output_height), scale, (min_ratio, max_ratio));
let descriptor = unsafe { MTLFXTemporalScalerDescriptor::new() };
unsafe {
descriptor.setColorTextureFormat(MTLPixelFormat::RGBA16Float);
descriptor.setDepthTextureFormat(MTLPixelFormat::Depth32Float);
descriptor.setMotionTextureFormat(MTLPixelFormat::RG16Float);
descriptor.setOutputTextureFormat(MTLPixelFormat::RGBA16Float);
descriptor.setInputWidth(input_width as usize);
descriptor.setInputHeight(input_height as usize);
descriptor.setOutputWidth(output_width as usize);
descriptor.setOutputHeight(output_height as usize);
descriptor.setAutoExposureEnabled(false);
}
let reactive_available = objc2::runtime::NSObjectProtocol::respondsToSelector(
&*descriptor,
objc2::sel!(setReactiveMaskTextureEnabled:),
);
if reactive_available {
unsafe {
descriptor.setReactiveMaskTextureEnabled(true);
descriptor.setReactiveMaskTextureFormat(REACTIVE_MASK_FORMAT);
}
}
let scaler =
unsafe { descriptor.newTemporalScalerWithDevice(device) }.ok_or_else(|| {
RenderError::Other("MetalFX: failed to create temporal scaler".to_string())
})?;
let required_output_usage = unsafe { scaler.outputTextureUsage() };
let output_desc = TextureDesc {
format: MTLPixelFormat::RGBA16Float,
width: output_width.max(1) as usize,
height: output_height.max(1) as usize,
usage: MTLTextureUsage(
required_output_usage.0
| MTLTextureUsage::ShaderRead.0
| MTLTextureUsage::RenderTarget.0,
),
..Default::default()
}
.build();
let output = device
.newTextureWithDescriptor(&output_desc)
.ok_or_else(|| allocation_failed("MetalFX upscaler output texture"))?;
let reactive = reactive_available && {
let needed = unsafe { scaler.reactiveTextureUsage() };
needed.0 & !REACTIVE_MASK_USAGE.0 == 0
};
if reactive_available && !reactive {
tracing::warn!(
"MetalFX: the reactive mask lacks a usage the scaler needs; running without it"
);
}
Ok(MetalFXUpscaler {
scaler,
output,
input_width,
input_height,
output_width,
output_height,
reactive,
})
}
}
impl MtlContext {
pub(in crate::metal) fn encode_upscale(
&self,
cmd_buf: &ProtocolObject<dyn objc2_metal::MTLCommandBuffer>,
scene_pre_taa: &Retained<ProtocolObject<dyn MTLTexture>>,
reactive_written: bool,
) -> RenderResult<u32> {
let upscaler = self.upscale.scaler.as_ref().ok_or_else(|| {
RenderError::Other("Upscale enabled but upscaler missing".to_string())
})?;
let gbuf = self.gbuffer.targets.as_ref().ok_or_else(|| {
RenderError::Other("Upscale enabled but G-buffer targets missing".to_string())
})?;
let velocity = self.gbuffer_velocity().ok_or_else(|| {
RenderError::Other(
"Upscale enabled but the pooled G-buffer velocity is missing".to_string(),
)
})?;
unsafe {
upscaler
.scaler
.setColorTexture(Some(scene_pre_taa.as_ref()));
upscaler.scaler.setDepthTexture(Some(gbuf.depth.as_ref()));
upscaler.scaler.setMotionTexture(Some(velocity));
upscaler
.scaler
.setOutputTexture(Some(upscaler.output.as_ref()));
if upscaler.reactive {
let mask = reactive_written.then(|| self.targets.hdr.reactive_mask.as_ref());
upscaler.scaler.setReactiveMaskTexture(mask);
}
upscaler
.scaler
.setMotionVectorScaleX(upscaler.input_width as f32);
upscaler
.scaler
.setMotionVectorScaleY(upscaler.input_height as f32);
let [jx, jy] = self
.upscale
.jitter
.load(std::sync::atomic::Ordering::Relaxed);
upscaler.scaler.setJitterOffsetX(jx);
upscaler.scaler.setJitterOffsetY(jy);
upscaler.scaler.setDepthReversed(CAMERA_DEPTH.reversed);
upscaler
.scaler
.setReset(crate::upscale_reset::consume(&self.upscale.reset));
}
unsafe {
upscaler.scaler.encodeToCommandBuffer(cmd_buf);
}
Ok(0)
}
}
#[derive(Default)]
pub(crate) struct UpscaleJitter(std::sync::atomic::AtomicU64);
impl UpscaleJitter {
pub(crate) fn store(&self, jx: f32, jy: f32, ordering: std::sync::atomic::Ordering) {
let packed = (jx.to_bits() as u64) | ((jy.to_bits() as u64) << 32);
self.0.store(packed, ordering);
}
pub(crate) fn load(&self, ordering: std::sync::atomic::Ordering) -> [f32; 2] {
let packed = self.0.load(ordering);
let jx = f32::from_bits(packed as u32);
let jy = f32::from_bits((packed >> 32) as u32);
[jx, jy]
}
}
#[cfg(test)]
mod tests {
use super::*;
const RANGE: (f32, f32) = (1.0, 3.0);
#[test]
fn the_realized_scale_does_not_round_trip() {
let out = (2048, 1536);
let built = scaler_input_size(out, 2.0 / 3.0, RANGE);
assert_eq!(built, (1365, 1024));
let realized = built.0 as f32 / out.0 as f32;
let rederived = scaler_input_size(out, realized, RANGE);
assert_ne!(
rederived, built,
"if this ever round-trips the guard below is measuring nothing"
);
assert_eq!(rederived.1, 1023, "a row short of the built input");
}
#[test]
fn a_request_outside_the_device_range_is_clamped() {
let out = (1920, 1080);
assert_eq!(scaler_input_size(out, 1.0 / 5.0, RANGE), (640, 360));
assert_eq!(scaler_input_size(out, 2.0, RANGE), out);
}
#[test]
fn degenerate_inputs_stay_renderable() {
assert_eq!(scaler_input_size((800, 600), 0.0, RANGE), (800, 600));
assert_eq!(scaler_input_size((800, 600), -1.0, RANGE), (800, 600));
assert_eq!(scaler_input_size((1, 1), 1.0 / 3.0, RANGE), (1, 1));
assert_eq!(scaler_input_size((1920, 1080), 0.5, (3.0, 1.0)), (640, 360));
}
}