-
Notifications
You must be signed in to change notification settings - Fork 4
Expand file tree
/
Copy pathserver.py
More file actions
99 lines (67 loc) · 2.1 KB
/
Copy pathserver.py
File metadata and controls
99 lines (67 loc) · 2.1 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
import argparse
def argparser():
parser = argparse.ArgumentParser(description='Hyperparameters')
parser.add_argument('--dataset', metavar='D', type=str, default="samsum",
help='Dataset to use')
parser.add_argument('--port', metavar='P', type=int, default=30327,
help='Port number')
# TODO: select bot (chatbot, sqlbot, etc.)
args = parser.parse_args()
return args
import os # nopep8
import sys # nopep8
sys.path.append(os.path.dirname(os.path.abspath(__file__))+'/bot') # nopep8
sys.path.append(os.path.dirname(os.path.abspath(__file__))+'/net') # nopep8
"""flask"""
from flask import Flask, request, jsonify # nopep8
from flask_cors import CORS # nopep8
app = Flask(__name__)
CORS(app)
"""bot"""
from bot.chatbot import Chatbot # nopep8
from net.gptj_lora import gptj_lora # nopep8
config, tokenizer, model = gptj_lora(
"models/36eca1e38b0d04afd013a735f4af49f77c15fbb1e93167ddd083b1548b66ab0a"
)
chatbot = Chatbot(tokenizer, model)
"""init"""
from datasets import load_dataset # nopep8
"""server"""
@app.route('/chat', methods=['POST'])
def chat():
params = request.get_json()
result = chatbot(**params)
return jsonify(dict({
'result': result
}))
@app.route('/train', methods=['POST'])
def train():
params = request.get_json()
params['dataset'] = load_dataset(
"samsum",
split=f"train[{params['from']}:{params['to']}]"
)
prev_version = chatbot.version
curr_version = chatbot.train(**params)
return jsonify(dict({
'previous': prev_version,
'current': curr_version
}))
@app.route('/aggregate', methods=['POST'])
def aggregate():
params = request.get_json()
prev_version = chatbot.version
curr_version = chatbot.aggregate(**params)
return jsonify(dict({
'previous': prev_version,
'current': curr_version
}))
"""main"""
if __name__ == "__main__":
# argparse
args = argparser()
print(args)
# init
train_dataset = load_dataset(args.dataset)
# run
app.run(host='0.0.0.0', port=args.port)