FZGPUModules 2.0
GPU-accelerated modular compression pipelines
Loading...
Searching...
No Matches
stage.h
Go to the documentation of this file.
1
5#pragma once
6
7#include "backend/types.h"
8#include "fzm_format.h"
9#include "stage/fusion.h"
10#include <array>
11#include <cstdint>
12#include <stdexcept>
13#include <string>
14#include <unordered_map>
15#include <vector>
16
17namespace fz {
18
19// Forward declaration — avoids requiring mempool.h in every stage header
20class MemoryPool;
21
31class Stage {
32public:
33 virtual ~Stage() = default;
34
52 static constexpr bool isSupportedOnBackend() { return true; }
53
68 virtual void execute(
69 fz::stream_t stream,
70 MemoryPool* pool,
71 const std::vector<void*>& inputs,
72 const std::vector<void*>& outputs,
73 const std::vector<size_t>& sizes
74 ) = 0;
75
77 virtual std::string getName() const = 0;
78
79 virtual size_t getNumInputs() const = 0;
80 virtual size_t getNumOutputs() const = 0;
81
88 virtual size_t getRequiredInputAlignment() const { return 1; }
89
94 virtual std::vector<std::string> getOutputNames() const {
95 return {"output"};
96 }
97
99 int getOutputIndex(const std::string& name) const {
100 auto names = getOutputNames();
101 for (size_t i = 0; i < names.size(); i++) {
102 if (names[i] == name) return static_cast<int>(i);
103 }
104 return -1;
105 }
106
112 virtual std::vector<size_t> estimateOutputSizes(
113 const std::vector<size_t>& input_sizes
114 ) const = 0;
115
117 virtual std::unordered_map<std::string, size_t> getActualOutputSizesByName() const = 0;
118
125 virtual size_t getActualOutputSize(int index) const {
126 auto names = getOutputNames();
127 if (index < 0 || index >= static_cast<int>(names.size())) return 0;
129 auto it = m.find(names[index]);
130 return (it != m.end()) ? it->second : 0;
131 }
132
137 virtual void setInverse(bool inverse) { (void)inverse; }
138 virtual bool isInverse() const { return false; }
139
141 virtual uint16_t getStageTypeId() const = 0;
142
144 virtual uint8_t getOutputDataType(size_t output_index) const = 0;
145
155 virtual uint8_t getInputDataType(size_t /*input_index*/) const {
156 return static_cast<uint8_t>(DataType::UNKNOWN);
157 }
158
163 virtual size_t serializeHeader(size_t output_index, uint8_t* header_buffer, size_t max_size) const {
164 (void)output_index; (void)header_buffer; (void)max_size;
165 return 0;
166 }
167
169 virtual void deserializeHeader(const uint8_t* header_buffer, size_t size) {
170 (void)header_buffer; (void)size;
171 }
172
179 virtual void saveState() {}
180 virtual void restoreState() {}
181
200 virtual std::vector<std::string> getRunNotes() const { return {}; }
201
207 virtual void setDims(const std::array<size_t, 3>& dims) { (void)dims; }
208
226 virtual void onFinalize(size_t /*estimated_inlen*/, MemoryPool* /*pool*/) {}
227
233 virtual size_t estimateDeviceFootprintBytes(size_t /*inlen*/) const { return 0; }
234
240 virtual size_t estimatePinnedFootprintBytes(size_t /*inlen*/) const { return 0; }
241
248 virtual void postStreamSync(fz::stream_t stream) { (void)stream; }
249
251 virtual size_t getMaxHeaderSize(size_t output_index) const {
252 (void)output_index;
253 return 0;
254 }
255
268 virtual bool isGraphCompatible() const { return true; }
269
278 virtual void setTerminalOutput(bool terminal) { (void)terminal; }
279
288 virtual size_t estimateScratchBytes(const std::vector<size_t>& input_sizes) const {
289 (void)input_sizes;
290 return 0;
291 }
292
303 virtual FusionSpec getFusionSpec() const { return {}; }
304
314 virtual FusedOpDecl getFusedOp() const { return {}; }
315
322 virtual EncodingOracleDecl getEncodingOracle() const { return {}; }
323
329 virtual bool bindDownstreamEncodingOracle(const EncodingOracleDecl& /*decl*/) {
330 return false;
331 }
332
334 virtual std::vector<FusedAuxOutputDecl> getFusedAuxOutputs() const { return {}; }
335
343 virtual void primeFusedForwardState(const FusedPrimeContext& /*ctx*/) {}
344
353 virtual void setFusedArchiveResult(size_t /*archive_bytes*/, size_t /*orig_bytes*/) {}
354
363 virtual void setFusedSideOutput(int /*output_index*/, size_t /*bytes*/) {}
364};
365
366} // namespace fz
Definition mempool.h:82
Definition stage.h:31
virtual size_t estimateDeviceFootprintBytes(size_t) const
Definition stage.h:233
virtual void setTerminalOutput(bool terminal)
Definition stage.h:278
virtual uint8_t getInputDataType(size_t) const
Definition stage.h:155
virtual std::string getName() const =0
virtual std::vector< std::string > getOutputNames() const
Definition stage.h:94
virtual std::vector< FusedAuxOutputDecl > getFusedAuxOutputs() const
Definition stage.h:334
virtual bool bindDownstreamEncodingOracle(const EncodingOracleDecl &)
Definition stage.h:329
virtual size_t estimatePinnedFootprintBytes(size_t) const
Definition stage.h:240
virtual void saveState()
Definition stage.h:179
virtual FusedOpDecl getFusedOp() const
Definition stage.h:314
virtual void primeFusedForwardState(const FusedPrimeContext &)
Definition stage.h:343
virtual bool isGraphCompatible() const
Definition stage.h:268
virtual void setFusedArchiveResult(size_t, size_t)
Definition stage.h:353
virtual size_t getActualOutputSize(int index) const
Definition stage.h:125
virtual void onFinalize(size_t, MemoryPool *)
Definition stage.h:226
virtual std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const =0
virtual void setInverse(bool inverse)
Definition stage.h:137
virtual uint16_t getStageTypeId() const =0
virtual void setFusedSideOutput(int, size_t)
Definition stage.h:363
virtual size_t serializeHeader(size_t output_index, uint8_t *header_buffer, size_t max_size) const
Definition stage.h:163
virtual uint8_t getOutputDataType(size_t output_index) const =0
virtual void deserializeHeader(const uint8_t *header_buffer, size_t size)
Definition stage.h:169
virtual size_t getRequiredInputAlignment() const
Definition stage.h:88
virtual std::unordered_map< std::string, size_t > getActualOutputSizesByName() const =0
virtual void execute(fz::stream_t stream, MemoryPool *pool, const std::vector< void * > &inputs, const std::vector< void * > &outputs, const std::vector< size_t > &sizes)=0
virtual void postStreamSync(fz::stream_t stream)
Definition stage.h:248
virtual EncodingOracleDecl getEncodingOracle() const
Definition stage.h:322
int getOutputIndex(const std::string &name) const
Definition stage.h:99
virtual size_t estimateScratchBytes(const std::vector< size_t > &input_sizes) const
Definition stage.h:288
virtual std::vector< std::string > getRunNotes() const
Definition stage.h:200
virtual FusionSpec getFusionSpec() const
Definition stage.h:303
static constexpr bool isSupportedOnBackend()
Definition stage.h:52
virtual size_t getMaxHeaderSize(size_t output_index) const
Definition stage.h:251
virtual void setDims(const std::array< size_t, 3 > &dims)
Definition stage.h:207
FZM binary file format definitions — structs, enums, and helpers.
Definition dag.h:24
@ UNKNOWN
Byte-transparent stages: skip type checking at finalize()
Host-side declaration of a local, exact encoded-size oracle.
Definition fusion.h:104
A stage's contribution to a generated fused kernel — the device-op it maps to, where its source lives...
Definition fusion.h:161
Minimal context a fused runner hands a stage so it can establish the forward-computed state its OWN i...
Definition fusion.h:186
A stage's fusion contract. Stages that can participate in a fused kernel override Stage::getFusionSpe...
Definition fusion.h:52
Backend-neutral GPU type aliases.