Skip to main content

lumen_engine/node/processing/
curves.rs

1use crate::node::{NodeId, NodeParamEvalContext, NodeParams, PortRef};
2
3use crate::gpu::{
4    BoundFrame, CompiledOutput, FrameBindContext, GpuCompileNode, GpuCompiledNode, RasterHandle,
5    compiler,
6};
7
8pub(crate) const SHADER: &str = include_str!("curves.wgsl");
9
10/// Applies a 1D RGB curve table to a raster.
11#[derive(Debug, Clone, lumen_macros::Delegate)]
12pub struct CurvesParams {
13    /// Curve table data source or named curve preset.
14    #[meta(name = "Curve", role = "curve_source", multiline, recommended_rows = 4)]
15    pub curve_source: String,
16    /// Blend amount for the curve adjustment.
17    #[meta(min = 0, max = 1, step = 0.01)]
18    pub strength: f64,
19}
20
21impl Default for CurvesParams {
22    fn default() -> Self {
23        Self {
24            curve_source: "identity".to_string(),
25            strength: 1.0,
26        }
27    }
28}
29
30/// Applies a 1D RGB curve table to a raster.
31#[derive(Debug, Clone, lumen_macros::Node)]
32#[node(kind = "curves", name = "Curves", category = "processing")]
33pub struct Curves {
34    pub id: NodeId,
35    #[params]
36    pub params: CurvesParamsDelegate,
37
38    #[input()]
39    pub source: PortRef,
40}
41
42impl Default for Curves {
43    fn default() -> Self {
44        Self {
45            id: NodeId::new(0),
46            params: CurvesParamsDelegate::default(),
47            source: PortRef::empty(),
48        }
49    }
50}
51
52impl GpuCompileNode for Curves {
53    fn compile_gpu(
54        &self,
55        ctx: &mut crate::gpu::CompileContext<'_>,
56        port: &PortRef,
57    ) -> crate::Result<CompiledOutput> {
58        if port.port != "output" {
59            return Err(ctx.missing_output(self.id, &port.port));
60        }
61
62        let source = ctx
63            .compile_port(&self.source)?
64            .into_raster(self.source.id, &self.source.port)?;
65        let size = source.domain.storage_size;
66        let texture = ctx.builder_mut().texture_for(
67            lumen_gpu::NodeKey(self.id.0),
68            Some(format!("curves:{}:output", self.id.0)),
69            lumen_gpu::TextureDesc::storage(size, lumen_gpu::wgpu::TextureFormat::Rgba8Unorm),
70        );
71        let params = ctx.builder_mut().buffer_for(
72            lumen_gpu::NodeKey(self.id.0),
73            Some(format!("curves:{}:params", self.id.0)),
74            lumen_gpu::BufferDesc::uniform(std::mem::size_of::<compiler::CurvesParams>() as u64),
75        );
76        let curve = ctx.builder_mut().buffer_for(
77            lumen_gpu::NodeKey(self.id.0),
78            Some(format!("curves:{}:table", self.id.0)),
79            lumen_gpu::BufferDesc::storage(std::mem::size_of::<compiler::CurvesTable>() as u64),
80        );
81        let program = ctx.builder_mut().program_for(
82            lumen_gpu::NodeKey(self.id.0),
83            lumen_gpu::ProgramDesc::Compute(lumen_gpu::ComputeProgramDesc {
84                label: Some("curves".to_string()),
85                shader: SHADER.to_string(),
86                entry: "cs_main".to_string(),
87                bind_groups: lumen_gpu::BindGroupLayoutSpec::single(vec![
88                    lumen_gpu::BindingLayoutEntry::texture(
89                        0,
90                        lumen_gpu::wgpu::ShaderStages::COMPUTE,
91                    ),
92                    lumen_gpu::BindingLayoutEntry::uniform(
93                        1,
94                        lumen_gpu::wgpu::ShaderStages::COMPUTE,
95                    ),
96                    lumen_gpu::BindingLayoutEntry::storage(
97                        2,
98                        lumen_gpu::wgpu::ShaderStages::COMPUTE,
99                        true,
100                    ),
101                    lumen_gpu::BindingLayoutEntry::storage_texture(
102                        3,
103                        lumen_gpu::wgpu::ShaderStages::COMPUTE,
104                        lumen_gpu::wgpu::TextureFormat::Rgba8Unorm,
105                        lumen_gpu::wgpu::StorageTextureAccess::WriteOnly,
106                    ),
107                ]),
108            }),
109        );
110        ctx.builder_mut().compute_pass(lumen_gpu::ComputePassDesc {
111            label: Some(format!("curves:{}:apply", self.id.0)),
112            owner: Some(lumen_gpu::NodeKey(self.id.0)),
113            program,
114            bindings: vec![
115                lumen_gpu::Binding::sampled_texture(0, 0, source.texture),
116                lumen_gpu::Binding::uniform(0, 1, params),
117                lumen_gpu::Binding::storage_buffer(0, 2, curve),
118                lumen_gpu::Binding::storage_texture(0, 3, texture),
119            ],
120            dispatch: compiler::dispatch_for(size).into(),
121        });
122        ctx.builder_mut().param(
123            lumen_gpu::ParamKey {
124                owner: lumen_gpu::NodeKey(self.id.0),
125                slot: 0,
126            },
127            lumen_gpu::ParamTarget::Buffer(params),
128        );
129        ctx.builder_mut().param(
130            lumen_gpu::ParamKey {
131                owner: lumen_gpu::NodeKey(self.id.0),
132                slot: 1,
133            },
134            lumen_gpu::ParamTarget::Buffer(curve),
135        );
136        ctx.register_compiled_node(CompiledCurves {
137            node_id: self.id,
138            params: self.params.clone(),
139            params_buffer: params,
140            curve_buffer: curve,
141        });
142
143        Ok(CompiledOutput::Raster(RasterHandle {
144            texture,
145            domain: source.domain,
146            metadata: source.metadata,
147        }))
148    }
149}
150
151#[derive(Debug, Clone)]
152struct CompiledCurves {
153    node_id: NodeId,
154    params: CurvesParamsDelegate,
155    params_buffer: lumen_gpu::BufferId,
156    curve_buffer: lumen_gpu::BufferId,
157}
158
159impl GpuCompiledNode for CompiledCurves {
160    fn node_id(&self) -> NodeId {
161        self.node_id
162    }
163
164    fn bind(&self, ctx: &FrameBindContext<'_>, bound: &mut BoundFrame) -> crate::Result<()> {
165        let evaluated = self.params.eval(&NodeParamEvalContext {
166            node_id: self.node_id,
167            expr: &ctx.expr_context(self.node_id, "params"),
168        })?;
169        let gpu_params = compiler::CurvesParams {
170            values: [evaluated.strength as f32, 0.0, 0.0, 0.0],
171        };
172        let curve =
173            compiler::CurvesTable::parse(self.node_id, ctx.frame(), &evaluated.curve_source)?;
174        bound.write_buffer(self.params_buffer, 0, bytemuck::bytes_of(&gpu_params));
175        bound.write_buffer(self.curve_buffer, 0, bytemuck::bytes_of(&curve));
176        Ok(())
177    }
178}