OpenDILabCommunity/DI-sheep
7
1import time2import numpy as np3import torch4from flask import Flask, request, jsonify, make_response5from flask_restplus import Api, Resource, fields6from threading import Thread7from sheep_env import SheepEnv8from sheep_model import SheepModel9 10flask_app = Flask(__name__)11app = Api(12 app=flask_app,13 version="0.0.1",14 title="DI-sheep App",15 description="Play Sheep with Deep Reinforcement Learning, Powered by OpenDILab"16)17 18name_space = app.namespace('DI-sheep', description='DI-sheep APIs')19model = app.model(20 'DI-sheep params', {21 'command': fields.String(required=False, description="Command Field", help="reset, step"),22 'argument': fields.Integer(required=False, description="Argument Field", help="reset->level, step->action"),23 }24)25MAX_ENV_NUM = 5026ENV_TIMEOUT_SECOND = 6027envs = {}28model = SheepModel(item_obs_size=80, item_num=30, global_obs_size=19)29ckpt = torch.load('ckpt_best.pth.tar', map_location='cpu')['model']30ckpt = {'item_encoder.encoder' + k.split('item_encoder')[-1] if 'item_encoder' in k else k: v for k, v in ckpt.items()} # compatibility for v1 and v2 model31model.load_state_dict(ckpt)32 33 34def random_action(obs, env):35 action_mask = obs['action_mask']36 action = np.random.choice(len(action_mask), p=action_mask / action_mask.sum())37 return action38 39 40def env_monitor():41 while True:42 cur_time = time.time()43 pop_keys = []44 for k, v in envs.items():45 if cur_time - v['update_time'] >= ENV_TIMEOUT_SECOND:46 pop_keys.append(k)47 for k in pop_keys:48 envs.pop(k)49 time.sleep(1)50 51 52app.env_thread = Thread(target=env_monitor, daemon=True)53app.env_thread.start()54 55 56@name_space.route("/")57class MainClass(Resource):58 59 def options(self):60 response = make_response()61 response.headers.add("Access-Control-Allow-Origin", "*")62 response.headers.add('Access-Control-Allow-Headers', "*")63 response.headers.add('Access-Control-Allow-Methods', "*")64 return response65 66 @app.expect(model)67 def post(self):68 try:69 t_start = time.time()70 data = request.json71 cmd, arg, uid = data['command'], data['argument'], data['uid']72 ip = request.remote_addr + uid73 74 if ip not in envs:75 if cmd == 'reset':76 if len(envs) >= MAX_ENV_NUM:77 response = jsonify(78 {79 "statusCode": 501,80 "status": "No enough env resource, please wait a moment",81 }82 )83 response.headers.add('Access-Control-Allow-Origin', '*')84 return response85 else:86 env = SheepEnv(1, agent=True, max_padding=True)87 env.seed(0)88 envs[ip] = {'env': env, 'update_time': time.time()}89 else:90 response = jsonify(91 {92 "statusCode": 501,93 "status": "No response for too long time, please reset the game",94 }95 )96 response.headers.add('Access-Control-Allow-Origin', '*')97 return response98 else:99 env = envs[ip]['env']100 envs[ip]['update_time'] = time.time()101 102 if cmd == 'reset':103 obs = env.reset(arg)104 action = model.compute_action(obs)105 # action = random_action(obs, env)106 scene = [item.to_json() for item in env.scene if item is not None]107 response = jsonify(108 {109 "statusCode": 200,110 "status": "Execution action",111 "result": {112 "scene": scene,113 "max_item_num": env.total_item_num,114 "action": action,115 }116 }117 )118 elif cmd == 'step':119 obs, _, done, _ = env.step(arg)120 action = model.compute_action(obs)121 # action = random_action(obs, env)122 scene = [item.to_json() for item in env.scene if item is not None]123 bucket = [item.to_json() for item in env.bucket]124 response = jsonify(125 {126 "statusCode": 200,127 "status": "Execution action",128 "result": {129 "scene": scene,130 "bucket": bucket,131 "done": done,132 "action": action,133 }134 }135 )136 else:137 response = jsonify({138 "statusCode": 500,139 "status": "Invalid command: {}".format(cmd),140 })141 response.headers.add('Access-Control-Allow-Origin', '*')142 return response143 print('backend process time: {}'.format(time.time() - t_start))144 print('current env number: {}'.format(len(envs)))145 response.headers.add('Access-Control-Allow-Origin', '*')146 return response147 except Exception as e:148 import traceback149 print(repr(e))150 print(traceback.format_exc())151 response = jsonify({152 "statusCode": 500,153 "status": "Could not execute action",154 })155 response.headers.add('Access-Control-Allow-Origin', '*')156 return response157 158if __name__ == "__main__":159 flask_app.run()160 161 162 