mirror of
https://github.com/wassname/stampy-chat.git
synced 2026-09-11 12:50:34 +08:00
working SSE streaming
This commit is contained in:
+1
-1
@@ -1 +1 @@
|
||||
web: gunicorn main:app
|
||||
web: gunicorn main:app --worker-class eventlet --threads 4
|
||||
|
||||
+52
-41
@@ -24,7 +24,7 @@ CONTEXT_FRACTION = 0.5 # the (approximate) fraction of num_tokens to use for co
|
||||
|
||||
ENCODER = tiktoken.get_encoding("cl100k_base")
|
||||
|
||||
DEBUG_PRINT = False
|
||||
DEBUG_PRINT = True
|
||||
|
||||
# --------------------------------- prompt code --------------------------------
|
||||
|
||||
@@ -116,48 +116,59 @@ def construct_prompt(query: str, history: List[Dict[str, str]], context: List[Bl
|
||||
return prompt
|
||||
|
||||
# ------------------------------- completion code -------------------------------
|
||||
import time
|
||||
import json
|
||||
|
||||
# returns either (True, reply string, top_k_blocks)) or (False, error message string, None)
|
||||
def talk_to_robot(index, query: str, history: List[Dict[str, str]], k: int = STANDARD_K):
|
||||
yield json.dumps({"state": "loading", "phase": "semantic"})
|
||||
time.sleep(1)
|
||||
yield json.dumps({"state": "loading", "phase": "prompt"})
|
||||
time.sleep(1)
|
||||
yield json.dumps({"state": "loading", "phase": "llm"})
|
||||
time.sleep(1)
|
||||
for c in "Hi. I'm a big dumb LLM. Hubris will be the end of us all.":
|
||||
yield json.dumps({"state": "streaming", "response": c})
|
||||
time.sleep(0.1)
|
||||
|
||||
try:
|
||||
# 1. Find the most relevant blocks from the Alignment Research Dataset
|
||||
top_k_blocks = get_top_k_blocks(index, query, k)
|
||||
|
||||
|
||||
# 2. Generate a prompt
|
||||
prompt = construct_prompt(query, history, top_k_blocks)
|
||||
|
||||
|
||||
# 3. Count number of tokens left for completion (-50 for a buffer)
|
||||
max_tokens_completion = NUM_TOKENS - sum([len(ENCODER.encode(message["content"]) + ENCODER.encode(message["role"])) for message in prompt]) - 50
|
||||
|
||||
# 4. Answer the user query
|
||||
t1 = time.time()
|
||||
response = openai.ChatCompletion.create(
|
||||
model=COMPLETIONS_MODEL,
|
||||
messages=prompt,
|
||||
max_tokens=max_tokens_completion
|
||||
)["choices"][0]["message"]["content"]
|
||||
t2 = time.time()
|
||||
print("Time to get response: ", t2 - t1)
|
||||
|
||||
|
||||
if DEBUG_PRINT:
|
||||
print('\n' * 10)
|
||||
print(" ------------------------------ prompt: -----------------------------")
|
||||
for message in prompt:
|
||||
print(f"----------- {message['role']}: ------------------")
|
||||
print(message['content'])
|
||||
|
||||
print('\n' * 10)
|
||||
|
||||
print(" ------------------------------ response: -----------------------------")
|
||||
print(response)
|
||||
|
||||
return (True, response, top_k_blocks)
|
||||
|
||||
except Exception as e:
|
||||
print(e)
|
||||
return (False, "Error: " + str(e), None)
|
||||
# try:
|
||||
# # 1. Find the most relevant blocks from the Alignment Research Dataset
|
||||
# top_k_blocks = get_top_k_blocks(index, query, k)
|
||||
#
|
||||
#
|
||||
# # 2. Generate a prompt
|
||||
# prompt = construct_prompt(query, history, top_k_blocks)
|
||||
#
|
||||
#
|
||||
# # 3. Count number of tokens left for completion (-50 for a buffer)
|
||||
# max_tokens_completion = NUM_TOKENS - sum([len(ENCODER.encode(message["content"]) + ENCODER.encode(message["role"])) for message in prompt]) - 50
|
||||
#
|
||||
# # 4. Answer the user query
|
||||
# t1 = time.time()
|
||||
# response = openai.ChatCompletion.create(
|
||||
# model=COMPLETIONS_MODEL,
|
||||
# messages=prompt,
|
||||
# max_tokens=max_tokens_completion
|
||||
# )["choices"][0]["message"]["content"]
|
||||
# t2 = time.time()
|
||||
# print("Time to get response: ", t2 - t1)
|
||||
#
|
||||
#
|
||||
# if DEBUG_PRINT:
|
||||
# print('\n' * 10)
|
||||
# print(" ------------------------------ prompt: -----------------------------")
|
||||
# for message in prompt:
|
||||
# print(f"----------- {message['role']}: ------------------")
|
||||
# print(message['content'])
|
||||
#
|
||||
# print('\n' * 10)
|
||||
#
|
||||
# print(" ------------------------------ response: -----------------------------")
|
||||
# print(response)
|
||||
#
|
||||
# return (True, response, top_k_blocks)
|
||||
#
|
||||
# except Exception as e:
|
||||
# print(e)
|
||||
# return (False, "Error: " + str(e), None)
|
||||
|
||||
|
||||
+19
-11
@@ -1,4 +1,4 @@
|
||||
from flask import Flask, jsonify, request
|
||||
from flask import Flask, jsonify, request, Response
|
||||
from flask_cors import CORS, cross_origin
|
||||
from get_blocks import get_top_k_blocks
|
||||
from chat import talk_to_robot
|
||||
@@ -34,6 +34,12 @@ app = Flask(__name__)
|
||||
cors = CORS(app)
|
||||
app.config['CORS_HEADERS'] = 'Content-Type'
|
||||
|
||||
# ---------------------------------- sse stuff ---------------------------------
|
||||
|
||||
def stream(src):
|
||||
yield from ('data: ' + '\ndata: '.join(message.splitlines()) + '\n\n' for message in src)
|
||||
yield 'data: close\n\n'
|
||||
|
||||
# ------------------------------- semantic search ------------------------------
|
||||
|
||||
|
||||
@@ -50,20 +56,22 @@ def semantic():
|
||||
@app.route('/chat', methods=['POST'])
|
||||
@cross_origin()
|
||||
def chat():
|
||||
#
|
||||
# query = request.json['query']
|
||||
# history = request.json['history']
|
||||
|
||||
return Response(stream(talk_to_robot(None, None, None)), mimetype='text/event-stream')
|
||||
|
||||
# is_valid, response, context = talk_to_robot(index, query, history)
|
||||
|
||||
query = request.json['query']
|
||||
history = request.json['history']
|
||||
|
||||
is_valid, response, context = talk_to_robot(index, query, history)
|
||||
|
||||
if is_valid:
|
||||
return jsonify({'response': response, 'citations': [{'title': block.title, 'author': block.author, 'date': block.date, 'url': block.url} for block in context]})
|
||||
else:
|
||||
return jsonify({'error': response})
|
||||
# if is_valid:
|
||||
# return jsonify({'response': response, 'citations': [{'title': block.title, 'author': block.author, 'date': block.date, 'url': block.url} for block in context]})
|
||||
# else:
|
||||
# return jsonify({'error': response})
|
||||
|
||||
|
||||
# ------------------------------------------------------------------------------
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
app.run(debug=True, port=3000)
|
||||
app.run(debug=True, port=3000)
|
||||
|
||||
+131
-87
@@ -121,10 +121,22 @@ const ShowEntry: React.FC<{entry: Entry}> = ({entry}) => {
|
||||
);
|
||||
};
|
||||
|
||||
type State = {
|
||||
state: "idle";
|
||||
} | {
|
||||
state: "loading";
|
||||
phase: "semantic" | "prompt" | "llm";
|
||||
} | {
|
||||
state: "streaming";
|
||||
response: string;
|
||||
};
|
||||
|
||||
|
||||
const Home: NextPage = () => {
|
||||
|
||||
const [ entries, setEntries ] = useState<Entry[]>([]);
|
||||
const [ runningIndex, setRunningIndex ] = useState(0);
|
||||
const [ loadState, setLoadState ] = useState<State>({state: "idle"});
|
||||
|
||||
const search = async (
|
||||
query: string,
|
||||
@@ -140,6 +152,37 @@ const Home: NextPage = () => {
|
||||
|
||||
setLoading(true);
|
||||
|
||||
|
||||
|
||||
// sse connection
|
||||
const eventSource = new EventSource(API_URL + "/chat");
|
||||
|
||||
eventSource.onmessage = (event) => {
|
||||
|
||||
if (event.data === "close") {
|
||||
eventSource.close();
|
||||
setLoading(false);
|
||||
setLoadState({state: "idle"});
|
||||
return;
|
||||
}
|
||||
|
||||
console.log(event.data);
|
||||
const data = JSON.parse(event.data);
|
||||
|
||||
if (data.state === "loading") {
|
||||
setLoadState({state: "loading", phase: data.phase});
|
||||
return;
|
||||
}
|
||||
|
||||
if (data.state === "streaming") {
|
||||
setLoadState((s) => {
|
||||
const response = s.state === "streaming" ? s.response : "";
|
||||
return {state: "streaming", response: response + data.response};
|
||||
});
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
const res = await fetch(API_URL + "/chat", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json", "Allow-Control-Allow-Origin": "*" },
|
||||
@@ -154,93 +197,93 @@ const Home: NextPage = () => {
|
||||
})
|
||||
})
|
||||
|
||||
if (!res.ok) {
|
||||
setLoading(false);
|
||||
console.log("load failure: " + res.status);
|
||||
return;
|
||||
}
|
||||
|
||||
const data = await res.json();
|
||||
|
||||
// -------------------------- error checking ---------------------------
|
||||
|
||||
if (data.error) {
|
||||
setEntries([...new_entries, {role: "error", content: data.error}]);
|
||||
setLoading(false);
|
||||
return;
|
||||
}
|
||||
|
||||
// ---------------------- normalize citation form ----------------------
|
||||
|
||||
// transform all things that look like [a, b, c] into [a][b][c]
|
||||
let response = data.response.replace(
|
||||
|
||||
/\[((?:[a-z]+,\s*)*[a-z]+)\]/g, // identify groups of this form
|
||||
|
||||
(block: string) => block.split(',')
|
||||
.map((x) => x.trim())
|
||||
.join("][")
|
||||
)
|
||||
|
||||
// transform all things that look like [(a), (b), (c)] into [(a)][(b)][(c)]
|
||||
response = response.replace(
|
||||
|
||||
/\[((?:\([a-z]+\),\s*)*\([a-z]+\))\]/g, // identify groups of this form
|
||||
|
||||
(block: string) => block.split(',')
|
||||
.map((x) => x.trim())
|
||||
.join("][")
|
||||
)
|
||||
|
||||
// transform all things that look like [(a)] into [a]
|
||||
response = response.replace(
|
||||
/\[\(([a-z]+)\)\]/g,
|
||||
(_match: string, x: string) => `[${x}]`
|
||||
)
|
||||
|
||||
// transform all things that look like [ a ] into [a]
|
||||
response = response.replace(
|
||||
/\[\s*([a-z]+)\s*\]/g,
|
||||
(_match: string, x: string) => `[${x}]`
|
||||
)
|
||||
|
||||
// -------------- map citations from strings into numbers --------------
|
||||
|
||||
// figure out what citations are in the response, and map them appropriately
|
||||
const cite_map = new Map<string, number>();
|
||||
let cite_count = runningIndex;
|
||||
|
||||
// scan a regex for [x] over the response. If x isn't in the map, add it.
|
||||
const regex = /\[([a-z]+)\]/g;
|
||||
let match;
|
||||
let response_copy = ""
|
||||
while ((match = regex.exec(response)) !== null) {
|
||||
if (!cite_map.has(match[1]!)) {
|
||||
cite_map.set(match[1]!, cite_count++);
|
||||
}
|
||||
// replace [x] with [i]
|
||||
response_copy += response.slice(response_copy.length, match.index) + `[${cite_map.get(match[1]!)! + 1}]`;
|
||||
}
|
||||
|
||||
setRunningIndex(cite_count);
|
||||
|
||||
response = response_copy + response.slice(response_copy.length);
|
||||
|
||||
// ----------------- create the ordered citation array -----------------
|
||||
|
||||
const citations = new Map<number, Citation>();
|
||||
cite_map.forEach((value, key) => {
|
||||
const index = key.charCodeAt(0) - 'a'.charCodeAt(0);
|
||||
if (index >= data.citations.length) {
|
||||
console.log("invalid citation index: " + index);
|
||||
} else {
|
||||
citations.set(value, data.citations[index]);
|
||||
}
|
||||
});
|
||||
|
||||
setEntries([...new_entries, {role: "assistant",
|
||||
content: response,
|
||||
citations: citations}]);
|
||||
// if (!res.ok) {
|
||||
// setLoading(false);
|
||||
// console.log("load failure: " + res.status);
|
||||
// return;
|
||||
// }
|
||||
//
|
||||
// const data = await res.json();
|
||||
//
|
||||
// // -------------------------- error checking ---------------------------
|
||||
//
|
||||
// if (data.error) {
|
||||
// setEntries([...new_entries, {role: "error", content: data.error}]);
|
||||
// setLoading(false);
|
||||
// return;
|
||||
// }
|
||||
//
|
||||
// // ---------------------- normalize citation form ----------------------
|
||||
//
|
||||
// // transform all things that look like [a, b, c] into [a][b][c]
|
||||
// let response = data.response.replace(
|
||||
//
|
||||
// /\[((?:[a-z]+,\s*)*[a-z]+)\]/g, // identify groups of this form
|
||||
//
|
||||
// (block: string) => block.split(',')
|
||||
// .map((x) => x.trim())
|
||||
// .join("][")
|
||||
// )
|
||||
//
|
||||
// // transform all things that look like [(a), (b), (c)] into [(a)][(b)][(c)]
|
||||
// response = response.replace(
|
||||
//
|
||||
// /\[((?:\([a-z]+\),\s*)*\([a-z]+\))\]/g, // identify groups of this form
|
||||
//
|
||||
// (block: string) => block.split(',')
|
||||
// .map((x) => x.trim())
|
||||
// .join("][")
|
||||
// )
|
||||
//
|
||||
// // transform all things that look like [(a)] into [a]
|
||||
// response = response.replace(
|
||||
// /\[\(([a-z]+)\)\]/g,
|
||||
// (_match: string, x: string) => `[${x}]`
|
||||
// )
|
||||
//
|
||||
// // transform all things that look like [ a ] into [a]
|
||||
// response = response.replace(
|
||||
// /\[\s*([a-z]+)\s*\]/g,
|
||||
// (_match: string, x: string) => `[${x}]`
|
||||
// )
|
||||
//
|
||||
// // -------------- map citations from strings into numbers --------------
|
||||
//
|
||||
// // figure out what citations are in the response, and map them appropriately
|
||||
// const cite_map = new Map<string, number>();
|
||||
// let cite_count = runningIndex;
|
||||
//
|
||||
// // scan a regex for [x] over the response. If x isn't in the map, add it.
|
||||
// const regex = /\[([a-z]+)\]/g;
|
||||
// let match;
|
||||
// let response_copy = ""
|
||||
// while ((match = regex.exec(response)) !== null) {
|
||||
// if (!cite_map.has(match[1]!)) {
|
||||
// cite_map.set(match[1]!, cite_count++);
|
||||
// }
|
||||
// // replace [x] with [i]
|
||||
// response_copy += response.slice(response_copy.length, match.index) + `[${cite_map.get(match[1]!)! + 1}]`;
|
||||
// }
|
||||
//
|
||||
// setRunningIndex(cite_count);
|
||||
//
|
||||
// response = response_copy + response.slice(response_copy.length);
|
||||
//
|
||||
// // ----------------- create the ordered citation array -----------------
|
||||
//
|
||||
// const citations = new Map<number, Citation>();
|
||||
// cite_map.forEach((value, key) => {
|
||||
// const index = key.charCodeAt(0) - 'a'.charCodeAt(0);
|
||||
// if (index >= data.citations.length) {
|
||||
// console.log("invalid citation index: " + index);
|
||||
// } else {
|
||||
// citations.set(value, data.citations[index]);
|
||||
// }
|
||||
// });
|
||||
//
|
||||
// setEntries([...new_entries, {role: "assistant",
|
||||
// content: response,
|
||||
// citations: citations}]);
|
||||
|
||||
setLoading(false);
|
||||
|
||||
@@ -261,6 +304,7 @@ const Home: NextPage = () => {
|
||||
))}
|
||||
</ul>
|
||||
<SearchBox search={search} />
|
||||
<p>{loadState.state === "loading" ? loadState.phase : ( loadState.state === "streaming" ? loadState.response : "" )}</p>
|
||||
</main>
|
||||
</>
|
||||
);
|
||||
|
||||
@@ -20,7 +20,7 @@ const SearchBox: React.FC<{search: (
|
||||
if (!loading) inputRef.current?.focus();
|
||||
}, [loading]);
|
||||
|
||||
if (loading) return <p>loading...</p>;
|
||||
if (loading) return <></>;
|
||||
return (<>
|
||||
<form className="flex mb-2" onSubmit={async (e) => {
|
||||
e.preventDefault();
|
||||
|
||||
Reference in New Issue
Block a user