CoolFace
Apppublic

OpenDILabCommunity/DI-sheep

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
7likes
agent_app.py162 linesDownload Raw Back to service
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