#![deny(unsafe_op_in_unsafe_fn)]
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;
pub(crate) struct UpscaleState {
pub scaler: Option<MetalFXUpscaler>,
pub scale: f32,
pub jitter: UpscaleJitter,
pub reset_pending: std::sync::atomic::AtomicBool,
}
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) 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,
) -> Result<Self, String> {
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 scaler = unsafe { descriptor.newTemporalScalerWithDevice(device) }
.ok_or_else(|| "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("MetalFX: failed to create upscaler output texture")?;
Ok(MetalFXUpscaler {
scaler,
output,
input_width,
input_height,
output_width,
output_height,
})
}
}
impl MtlContext {
pub(in crate::metal) fn encode_upscale(
&self,
cmd_buf: &ProtocolObject<dyn objc2_metal::MTLCommandBuffer>,
scene_pre_taa: &Retained<ProtocolObject<dyn MTLTexture>>,
) -> Result<u32, String> {
let upscaler = self
.upscale
.scaler
.as_ref()
.ok_or("Upscale enabled but upscaler missing")?;
let gbuf = self
.gbuffer
.targets
.as_ref()
.ok_or("Upscale enabled but G-buffer targets missing")?;
let velocity = self
.gbuffer_velocity()
.ok_or("Upscale enabled but the pooled G-buffer velocity is missing")?;
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()));
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(false);
let reset = self
.upscale
.reset_pending
.swap(false, std::sync::atomic::Ordering::AcqRel);
upscaler.scaler.setReset(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_realised_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 realised = built.0 as f32 / out.0 as f32;
let rederived = scaler_input_size(out, realised, 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));
}
}