Ragged Batch Builder#

class IdentitySetScratch#

Public Functions

void reserve(size_t maxItems)#
void clear()#
bool insert(uint64_t value)#
class RaggedBatchBuilder#

Public Functions

explicit RaggedBatchBuilder(RaggedEngineContract contract)#
void reserve(RaggedExecutionBatch &batch)#
void buildInto(
ScheduledStep const &step,
RaggedExecutionBatch &batch
)#
void finalizeSubsetSelection(
ScheduledStep const &sourceStep,
std::vector<int32_t> const &keepStartOffsets,
std::vector<int32_t> const &concatenatedKeepIndices,
RaggedExecutionBatch &batch
) const#

Public Static Functions

static void validateRuntimeAdapterStep(
ScheduledStep const &step,
int32_t activeBatchSize,
std::vector<RequestId> const &requestIds,
std::vector<ResidentRef> const &residentRefs
)#
static void validateExecutionBatch(
RaggedExecutionBatch const &batch,
RaggedEngineContract const &contract
)#
struct RaggedEngineContract#

Public Members

TokenLayoutBackend backend = {TokenLayoutBackend::kEntryPaddedCompatibility}#
int32_t maxNumSequences = {0}#
int32_t maxQueryLength = {0}#
int32_t maxPhysicalTokens = {0}#
int32_t recurrentPoolRows = {0}#
bool mixedStepSupported = {false}#
struct RaggedStepShape#

Public Members

int32_t numSequences = {0}#
int32_t validTokens = {0}#
int32_t physicalTokens = {0}#
int32_t queryWidth = {0}#
int32_t numContextSequences = {0}#
int32_t numContextTokens = {0}#
int32_t numLogits = {0}#
struct CompletionSequence#

Public Members

RequestId requestId = {0}#
ResidentRef resident#
int32_t pastLength = {0}#
struct RaggedExecutionBatch#

Public Functions

void validateCommitSnapshot(
StepId completedStepId,
std::vector<CompletionSequence> const &current
) const#

Public Members

StepId stepId = {0}#
TokenLayoutBackend layout = {TokenLayoutBackend::kEntryPaddedCompatibility}#
RaggedStepShape shape#
Tensor hostTokenIds#
std::vector<SequenceIdentity> sequenceOrder#
std::vector<int32_t> positions#
std::vector<int32_t> queryStartOffsets#
std::vector<int32_t> queryLengths#
std::vector<int32_t> pastLengths#
std::vector<int32_t> attentionSequenceLengths#
std::vector<int32_t> stateIndices#
std::vector<int64_t> logitsIndices#
std::vector<int32_t> logitsToSequence#
std::vector<SequenceWork> sequenceWorks#