/* -*- Mode: C++; tab-width: 8; indent-tabs-mode: nil; c-basic-offset: 2 -*- */ /* This Source Code Form is subject to the terms of the Mozilla Public * License, v. 2.0. If a copy of the MPL was not distributed with this * file, You can obtain one at http://mozilla.org/MPL/2.0/. */ using mozilla::dom::TextGenerationRole from "mozilla/hwinference/TextGenerationEnums.h"; using mozilla::dom::TextGenerationFinishReason from "mozilla/hwinference/TextGenerationEnums.h"; using mozilla::dom::TextGenerationSamplerType from "mozilla/hwinference/TextGenerationEnums.h"; using mozilla::dom::TextGenerationKVCacheDtype from "mozilla/hwinference/TextGenerationEnums.h"; namespace mozilla { namespace hwinference { struct ChatMessage { TextGenerationRole role; nsCString content; }; struct LogitBias { int32_t token; float bias; }; struct Sampler { TextGenerationSamplerType type; int32_t topK; float topP; float temp; // Only the dist sampler reads this. uint32_t seed; LogitBias[] logitBias; }; // Appended to the generator's context; Clear() first for a fresh start. struct GenerateRequest { ChatMessage[] messages; uint32_t maxTokens; // Number of generated tokens batched into one Delta. uint32_t bufferLength; Sampler[] samplers; int32_t[] stopTokens; bool stopOnEndOfGenerationTokens; }; struct Timings { double prefillMs; double decodeMs; }; struct LoadSuccess { double loadMs; }; struct LoadError { nsCString message; }; union LoadResult { LoadSuccess; LoadError; }; struct Usage { uint32_t promptTokens; // Unicode code points in the templated prompt. uint32_t promptCharacters; uint32_t generatedTokens; Timings timings; }; struct ResourceSnapshot { // Cumulative generator-process CPU since it started. uint64_t cpuTimeMs; // Private physical bytes. uint64_t memoryBytes; }; struct ResourceUsage { ResourceSnapshot before; ResourceSnapshot after; }; // content holds the complete generated text, including every Delta sent. struct GenerateResult { nsCString content; TextGenerationFinishReason reason; Usage usage; ResourceUsage resources; }; struct GenerateError { nsCString message; }; union GenerateResponse { GenerateResult; GenerateError; }; struct TextGenerationOptions { uint32_t contextSize; // numThreads is the prefill/batch count, numThreadsDecoding the decode one. uint32_t numThreads; uint32_t numThreadsDecoding; uint32_t batchSize; uint32_t ubatchSize; TextGenerationKVCacheDtype kvCacheDtype; bool flashAttn; }; } // namespace hwinference } // namespace mozilla