Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
59 changes: 58 additions & 1 deletion include/ryzenai/inference_engine.h
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
#include <vector>
#include <memory>
#include <mutex>
#include <unordered_map>

// Forward declarations for ONNX Runtime GenAI
struct OgaModel;
Expand All @@ -16,6 +17,9 @@ struct OgaMultiModalProcessor;

namespace ryzenai {

// Single global multi-turn session until clients opt in via conversation_id.
inline constexpr const char* kDefaultConversationId = "default";

// Timing data returned from completion
struct CompletionTimingData {
int token_count = 0; // Number of generated tokens
Expand All @@ -41,7 +45,41 @@ class InferenceEngine {
StreamCallback callback);

// Apply chat template to messages
std::string applyChatTemplate(const std::string& messages_json, const std::string& tools_json = "");
std::string applyChatTemplate(const std::string& messages_json,
const std::string& tools_json = "",
bool add_generation_prompt = true);

// Multi-turn chat (non-tool): reuse one OgaGenerator per conversation_id and
// AppendTokens only the delta prompt on subsequent turns.
std::string completeMultiTurn(const std::string& conversation_id,
const std::string& messages_json,
const GenerationParams& params,
CompletionTimingData* out_timing = nullptr);

void streamMultiTurn(const std::string& conversation_id,
const std::string& messages_json,
const GenerationParams& params,
StreamCallback callback);

void resetMultiTurnSession(const std::string& conversation_id);

// Multi-turn VLM: reuse ChatSession KV cache; turn 1 may SetInputs (vision),
// later text-only turns append delta from last <|im_start|>user via jinja template.
std::string completeMultiTurnMultimodal(const std::string& conversation_id,
const std::string& messages_json,
const std::vector<std::string>& new_turn_images,
const std::vector<std::string>& all_images,
const std::string& tools_json,
const GenerationParams& params,
CompletionTimingData* out_timing = nullptr);

void streamMultiTurnMultimodal(const std::string& conversation_id,
const std::string& messages_json,
const std::vector<std::string>& new_turn_images,
const std::vector<std::string>& all_images,
const std::string& tools_json,
const GenerationParams& params,
StreamCallback callback);

// Apply the model's chat template strictly via OGA (jinja), without the
// text-only manual fallbacks. Required for multimodal models whose template
Expand Down Expand Up @@ -85,6 +123,25 @@ class InferenceEngine {
std::string resolveModelPath(const std::string& path);
std::vector<int32_t> truncatePrompt(const std::vector<int32_t>& input_ids);
bool validateModelDirectory(const std::string& path);

struct ChatSession {
std::unique_ptr<OgaGenerator> generator;
std::unique_ptr<OgaGeneratorParams> gen_params;
size_t turn_count = 0;
};

ChatSession& getOrCreateChatSession(const std::string& conversation_id);
void resetMultiTurnSessionLocked(const std::string& conversation_id);
// Extract the latest user turn from full prompt: suffix starting at last <|im_start|>user.
std::string extractDeltaPromptFromLastUser(const std::string& full_prompt) const;
void appendPromptText(OgaGenerator& generator, const std::string& text);
void configureGeneratorParams(OgaGeneratorParams& gen_params,
const GenerationParams& params,
int total_max_length) const;
std::string applyStopSequences(const std::string& text, const GenerationParams& params) const;
std::vector<int32_t> encodeText(const std::string& text) const;

std::unordered_map<std::string, ChatSession> chat_sessions_;

std::unique_ptr<OgaModel> model_;
std::unique_ptr<OgaTokenizer> tokenizer_;
Expand Down
8 changes: 6 additions & 2 deletions include/ryzenai/server.h
Original file line number Diff line number Diff line change
Expand Up @@ -33,9 +33,13 @@ class RyzenAIServer {
void handleCompletions(const httplib::Request& req, httplib::Response& res);
void handleChatCompletions(const httplib::Request& req, httplib::Response& res);
void handleMultimodalChat(const json& request_json, const ChatCompletionRequest& chat_req,
const std::vector<std::string>& images, httplib::Response& res);
const json& messages_array,
const std::vector<std::string>& new_turn_images,
const std::vector<std::string>& all_images,
httplib::Response& res);
void handleResponses(const httplib::Request& req, httplib::Response& res);

void handleSessionReset(const httplib::Request& req, httplib::Response& res);

// Helper methods
json createErrorResponse(const std::string& message, const std::string& type);
void setupCORS(httplib::Response& res);
Expand Down
1 change: 1 addition & 0 deletions include/ryzenai/types.h
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@ struct CompletionRequest {
// Chat completion request (OpenAI format)
struct ChatCompletionRequest {
std::vector<ChatMessage> messages;
std::string conversation_id; // reserved: parsed but not required; server uses default session
int max_tokens = 1500;
float temperature = 0.7f;
float top_p = 0.9f;
Expand Down
Loading