first version
This commit is contained in:
parent
ae82ae53b2
commit
139b42a734
412
maskplace/PPO2.py
Normal file
412
maskplace/PPO2.py
Normal file
@ -0,0 +1,412 @@
|
||||
import argparse
|
||||
import pickle
|
||||
from collections import namedtuple
|
||||
|
||||
import os
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = "0"
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
import gym
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.optim as optim
|
||||
from torch.distributions import Normal
|
||||
from torch.distributions import Categorical
|
||||
from torch.utils.data.sampler import BatchSampler, SubsetRandomSampler
|
||||
|
||||
import place_env
|
||||
import torchvision
|
||||
from place_db import PlaceDB
|
||||
import time
|
||||
from tqdm import tqdm
|
||||
import random
|
||||
from comp_res import comp_res
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
|
||||
# set device to cpu or cuda
|
||||
device = torch.device('cuda')
|
||||
|
||||
if(torch.cuda.is_available()):
|
||||
device = torch.device('cuda:0')
|
||||
torch.cuda.empty_cache()
|
||||
print("Device set to : " + str(torch.cuda.get_device_name(device)))
|
||||
else:
|
||||
print("Device set to : cpu")
|
||||
|
||||
# Parameters
|
||||
parser = argparse.ArgumentParser(description='Solve the Pendulum-v0 with PPO')
|
||||
parser.add_argument(
|
||||
'--gamma', type=float, default=0.95, metavar='G', help='discount factor (default: 0.9)')
|
||||
parser.add_argument('--seed', type=int, default=42, metavar='N', help='random seed (default: 0)')
|
||||
parser.add_argument('--disable_tqdm', type=int, default=1)
|
||||
parser.add_argument('--lr', type=float, default=2.5e-3)
|
||||
parser.add_argument(
|
||||
'--log-interval',
|
||||
type=int,
|
||||
default=10,
|
||||
metavar='N',
|
||||
help='interval between training status logs (default: 10)')
|
||||
parser.add_argument('--pnm', type=int, default=128)
|
||||
parser.add_argument('--benchmark', type=str, default='adaptec1')
|
||||
parser.add_argument('--soft_coefficient', type=float, default = 1)
|
||||
parser.add_argument('--batch_size', type=int, default=64)
|
||||
parser.add_argument('--is_test', action='store_true', default=False)
|
||||
parser.add_argument('--save_fig', action='store_true', default=False)
|
||||
args = parser.parse_args()
|
||||
writer = SummaryWriter('./tb_log')
|
||||
|
||||
benchmark = args.benchmark
|
||||
placedb = PlaceDB(benchmark)
|
||||
grid = 224
|
||||
placed_num_macro = args.pnm
|
||||
if args.pnm > placedb.node_cnt:
|
||||
placed_num_macro = placedb.node_cnt
|
||||
args.pnm = placed_num_macro
|
||||
env = gym.make('place_env-v0', placedb = placedb, placed_num_macro = placed_num_macro, grid = grid).unwrapped
|
||||
|
||||
num_emb_state = 64 + 2 + 1
|
||||
num_state = 1 + grid*grid*5 + 2
|
||||
|
||||
def seed_torch(seed=0):
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
os.environ['PYTHONHASHSEED'] = str(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
env.seed(seed)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
num_action = env.action_space.shape
|
||||
seed_torch(args.seed)
|
||||
|
||||
Transition = namedtuple('Transition',['state', 'action', 'reward', 'a_log_prob', 'next_state', 'reward_intrinsic'])
|
||||
TrainingRecord = namedtuple('TrainRecord',['episode', 'reward'])
|
||||
print("seed = {}".format(args.seed))
|
||||
print("lr = {}".format(args.lr))
|
||||
print("placed_num_macro = {}".format(args.pnm))
|
||||
|
||||
|
||||
class MyCNN(nn.Module):
|
||||
def __init__(self):
|
||||
super(MyCNN, self).__init__()
|
||||
self.cnn = nn.Sequential(
|
||||
nn.Conv2d(4, 8, 1),
|
||||
nn.ReLU(),
|
||||
nn.Conv2d(8, 8, 1),
|
||||
nn.ReLU(),
|
||||
nn.Conv2d(8, 1, 1),
|
||||
)
|
||||
def forward(self, x):
|
||||
return self.cnn(x)
|
||||
|
||||
|
||||
class MyCNNCoarse(nn.Module):
|
||||
def __init__(self, res_net):
|
||||
super(MyCNNCoarse, self).__init__()
|
||||
self.cnn = res_net.to(device)
|
||||
self.cnn.fc = torch.nn.Linear(512, 16*7*7)
|
||||
self.deconv = nn.Sequential(
|
||||
nn.ConvTranspose2d(16, 8, 3, stride=2, padding=1, output_padding = 1), #14
|
||||
nn.ReLU(),
|
||||
nn.ConvTranspose2d(8, 4, 3, stride=2, padding=1, output_padding = 1), #28
|
||||
nn.ReLU(),
|
||||
nn.ConvTranspose2d(4, 2, 3, stride=2, padding=1, output_padding = 1), #56
|
||||
nn.ReLU(),
|
||||
nn.ConvTranspose2d(2, 1, 3, stride=2, padding=1, output_padding = 1), #112
|
||||
nn.ReLU(),
|
||||
nn.ConvTranspose2d(1, 1, 3, stride=2, padding=1, output_padding = 1), #224
|
||||
)
|
||||
def forward(self, x):
|
||||
x = self.cnn(x).reshape(-1, 16, 7, 7)
|
||||
return self.deconv(x)
|
||||
|
||||
|
||||
class Actor(nn.Module):
|
||||
def __init__(self, cnn, gcn, cnn_coarse):
|
||||
super(Actor, self).__init__()
|
||||
self.fc1 = nn.Linear(num_emb_state, 512)
|
||||
self.fc2 = nn.Linear(512, 64)
|
||||
self.fc3 = nn.Linear(64, grid * grid)
|
||||
self.cnn = cnn
|
||||
self.cnn_coarse = cnn_coarse
|
||||
self.gcn = None
|
||||
self.softmax = nn.Softmax(dim=-1)
|
||||
self.merge = nn.Conv2d(2, 1, 1)
|
||||
|
||||
def forward(self, x, graph = None, cnn_res = None, gcn_res = None, graph_node = None):
|
||||
if not cnn_res:
|
||||
cnn_input = x[:, 1+grid*grid*1: 1+grid*grid*5].reshape(-1, 4, grid, grid)
|
||||
mask = x[:, 1+grid*grid*2: 1+grid*grid*3].reshape(-1, grid, grid)
|
||||
mask = mask.flatten(start_dim=1, end_dim=2)
|
||||
cnn_res = self.cnn(cnn_input)
|
||||
coarse_input = torch.cat((x[:, 1: 1+grid*grid*2].reshape(-1, 2, grid, grid),
|
||||
x[:, 1+grid*grid*3: 1+grid*grid*4].reshape(-1, 1, grid, grid)
|
||||
),dim= 1).reshape(-1, 3, grid, grid)
|
||||
cnn_coarse_res = self.cnn_coarse(coarse_input)
|
||||
cnn_res = self.merge(torch.cat((cnn_res, cnn_coarse_res), dim=1))
|
||||
net_img = x[:, 1+grid*grid: 1+grid*grid*2]
|
||||
net_img = net_img + x[:, 1+grid*grid*2: 1+grid*grid*3] * 10
|
||||
net_img_min = net_img.min() + args.soft_coefficient
|
||||
mask2 = net_img.le(net_img_min).logical_not().float()
|
||||
|
||||
x = cnn_res
|
||||
x = x.reshape(-1, grid * grid)
|
||||
x = torch.where(mask + mask2 >=1.0, -1.0e10, x.double())
|
||||
x = self.softmax(x)
|
||||
|
||||
return x, cnn_res, gcn_res
|
||||
|
||||
|
||||
class Critic(nn.Module):
|
||||
def __init__(self, cnn, gcn, cnn_coarse, res_net):
|
||||
super(Critic, self).__init__()
|
||||
self.fc1 = nn.Linear(64, 64)
|
||||
self.fc2 = nn.Linear(64, 64)
|
||||
self.state_value = nn.Linear(64, 1)
|
||||
self.pos_emb = nn.Embedding(1400, 64)
|
||||
self.cnn = cnn
|
||||
self.gcn = gcn
|
||||
def forward(self, x, graph = None, cnn_res = None, gcn_res = None, graph_node = None):
|
||||
x1 = F.relu(self.fc1(self.pos_emb(x[:, 0].long())))
|
||||
x2 = F.relu(self.fc2(x1))
|
||||
value = self.state_value(x2)
|
||||
return value
|
||||
|
||||
|
||||
class PPO():
|
||||
clip_param = 0.2
|
||||
max_grad_norm = 0.5
|
||||
ppo_epoch = 10
|
||||
if placed_num_macro:
|
||||
buffer_capacity = 10 * (placed_num_macro)
|
||||
else:
|
||||
buffer_capacity = 5120
|
||||
batch_size = args.batch_size
|
||||
print("batch_size = {}".format(batch_size))
|
||||
|
||||
def __init__(self):
|
||||
super(PPO, self).__init__()
|
||||
self.gcn = None
|
||||
self.resnet = torchvision.models.resnet18(pretrained=True)
|
||||
self.cnn = MyCNN().to(device)
|
||||
self.cnn_coarse = MyCNNCoarse(self.resnet).to(device)
|
||||
self.actor_net = Actor(cnn = self.cnn, gcn = self.gcn, cnn_coarse = self.cnn_coarse).float().to(device)
|
||||
self.critic_net = Critic(cnn = self.cnn, gcn = self.gcn, cnn_coarse = None, res_net = self.resnet).float().to(device)
|
||||
self.buffer = []
|
||||
self.counter = 0
|
||||
self.training_step = 0
|
||||
self.actor_optimizer = optim.Adam(self.actor_net.parameters(), args.lr)
|
||||
self.critic_net_optimizer = optim.Adam(self.critic_net.parameters(), args.lr)
|
||||
|
||||
def load_param(self, path):
|
||||
checkpoint = torch.load(path, map_location=torch.device(device))
|
||||
self.actor_net.load_state_dict(checkpoint['actor_net_dict'])
|
||||
self.critic_net.load_state_dict(checkpoint['critic_net_dict'])
|
||||
|
||||
def select_action(self, state):
|
||||
state = torch.from_numpy(state).float().to(device).unsqueeze(0)
|
||||
with torch.no_grad():
|
||||
action_probs, _, _ = self.actor_net(state)
|
||||
dist = Categorical(action_probs)
|
||||
action = dist.sample()
|
||||
action_log_prob = dist.log_prob(action)
|
||||
return action.item(), action_log_prob.item()
|
||||
|
||||
def get_value(self, state):
|
||||
state = torch.from_numpy(state)
|
||||
with torch.no_grad():
|
||||
value = self.critic_net(state)
|
||||
return value.item()
|
||||
|
||||
def save_param(self, running_reward):
|
||||
strftime = time.strftime("%Y-%m-%d-%H-%M-%S", time.localtime())
|
||||
if not os.path.exists("save_models"):
|
||||
os.mkdir("save_models")
|
||||
torch.save({"actor_net_dict": self.actor_net.state_dict(),
|
||||
"critic_net_dict": self.critic_net.state_dict()},
|
||||
"./save_models/net_dict-{}-{}-".format(benchmark, placed_num_macro)+strftime+"{}".format(int(running_reward))+".pkl")
|
||||
|
||||
def store_transition(self, transition):
|
||||
self.buffer.append(transition)
|
||||
self.counter+=1
|
||||
return self.counter % self.buffer_capacity == 0
|
||||
|
||||
def update(self):
|
||||
state = torch.tensor(np.array([t.state for t in self.buffer]), dtype=torch.float)
|
||||
action = torch.tensor(np.array([t.action for t in self.buffer]), dtype=torch.float).view(-1, 1).to(device)
|
||||
reward = torch.tensor(np.array([t.reward for t in self.buffer]), dtype=torch.float).view(-1, 1).to(device)
|
||||
old_action_log_prob = torch.tensor(np.array([t.a_log_prob for t in self.buffer]), dtype=torch.float).view(-1, 1).to(device)
|
||||
del self.buffer[:]
|
||||
target_list = []
|
||||
target = 0
|
||||
for i in range(reward.shape[0]-1, -1, -1):
|
||||
if state[i, 0] >= placed_num_macro - 1:
|
||||
target = 0
|
||||
r = reward[i, 0].item()
|
||||
target = r + args.gamma * target
|
||||
target_list.append(target)
|
||||
target_list.reverse()
|
||||
target_v_all = torch.tensor(np.array([t for t in target_list]), dtype=torch.float).view(-1, 1).to(device)
|
||||
|
||||
for _ in range(self.ppo_epoch): # iteration ppo_epoch
|
||||
for index in tqdm(BatchSampler(SubsetRandomSampler(range(self.buffer_capacity)), self.batch_size, True),
|
||||
disable = args.disable_tqdm):
|
||||
self.training_step +=1
|
||||
|
||||
action_probs, _, _ = self.actor_net(state[index].to(device))
|
||||
dist = Categorical(action_probs)
|
||||
action_log_prob = dist.log_prob(action[index].squeeze())
|
||||
ratio = torch.exp(action_log_prob - old_action_log_prob[index].squeeze())
|
||||
target_v = target_v_all[index]
|
||||
critic_net_output = self.critic_net(state[index].to(device))
|
||||
advantage = (target_v - critic_net_output).detach()
|
||||
|
||||
L1 = ratio * advantage.squeeze()
|
||||
L2 = torch.clamp(ratio, 1-self.clip_param, 1+self.clip_param) * advantage.squeeze()
|
||||
action_loss = -torch.min(L1, L2).mean() # MAX->MIN desent
|
||||
|
||||
self.actor_optimizer.zero_grad()
|
||||
action_loss.backward()
|
||||
nn.utils.clip_grad_norm_(self.actor_net.parameters(), self.max_grad_norm)
|
||||
self.actor_optimizer.step()
|
||||
|
||||
value_loss = F.smooth_l1_loss(self.critic_net(state[index].to(device)), target_v)
|
||||
self.critic_net_optimizer.zero_grad()
|
||||
value_loss.backward()
|
||||
nn.utils.clip_grad_norm_(self.critic_net.parameters(), self.max_grad_norm)
|
||||
self.critic_net_optimizer.step()
|
||||
|
||||
writer.add_scalar('action_loss', action_loss, self.training_step)
|
||||
writer.add_scalar('value_loss', value_loss, self.training_step)
|
||||
|
||||
|
||||
def save_placement(file_path, node_pos, ratio):
|
||||
fwrite = open(file_path, 'w')
|
||||
node_place = {}
|
||||
for node_name in node_pos:
|
||||
|
||||
x, y,_ , _ = node_pos[node_name]
|
||||
x = round(x * ratio + ratio)
|
||||
y = round(y * ratio + ratio)
|
||||
node_place[node_name] = (x, y)
|
||||
print("len node_place", len(node_place))
|
||||
for node_name in placedb.node_info:
|
||||
if node_name not in node_place:
|
||||
continue
|
||||
x, y = node_place[node_name]
|
||||
fwrite.write('{}\t{}\t{}\t:\tN /FIXED\n'.format(node_name, x, y))
|
||||
print(".pl has been saved to {}.".format(file_path))
|
||||
|
||||
|
||||
def main():
|
||||
|
||||
agent = PPO()
|
||||
strftime = time.strftime("%Y-%m-%d-%H-%M-%S", time.localtime())
|
||||
|
||||
training_records = []
|
||||
running_reward = -1000000
|
||||
|
||||
|
||||
log_file_name = "logs/log_"+ benchmark + "_" + strftime + "_seed_"+ str(args.seed) + "_pnm_" + str(args.pnm) + ".csv"
|
||||
if not os.path.exists("logs"):
|
||||
os.mkdir("logs")
|
||||
fwrite = open(log_file_name, "w")
|
||||
load_model_path = None
|
||||
|
||||
if load_model_path:
|
||||
agent.load_param(load_model_path)
|
||||
|
||||
best_reward = running_reward
|
||||
if args.is_test:
|
||||
torch.inference_mode()
|
||||
|
||||
for i_epoch in range(100000):
|
||||
score = 0
|
||||
raw_score = 0
|
||||
start = time.time()
|
||||
state = env.reset()
|
||||
|
||||
done = False
|
||||
while done is False:
|
||||
state_tmp = state.copy()
|
||||
action, action_log_prob = agent.select_action(state)
|
||||
|
||||
next_state, reward, done, info = env.step(action)
|
||||
assert next_state.shape == (num_state, )
|
||||
reward_intrinsic = 0
|
||||
if not args.is_test:
|
||||
trans = Transition(state_tmp, action, reward / 200.0, action_log_prob, next_state, reward_intrinsic)
|
||||
if not args.is_test and agent.store_transition(trans):
|
||||
assert done == True
|
||||
agent.update()
|
||||
score += reward
|
||||
raw_score += info["raw_reward"]
|
||||
state = next_state
|
||||
end = time.time()
|
||||
|
||||
if i_epoch == 0:
|
||||
running_reward = score
|
||||
running_reward = running_reward * 0.9 + score * 0.1
|
||||
print("score = {}, raw_score = {}".format(score, raw_score))
|
||||
|
||||
if running_reward > best_reward * 0.975:
|
||||
best_reward = running_reward
|
||||
if i_epoch >= 10:
|
||||
agent.save_param(running_reward)
|
||||
if args.save_fig:
|
||||
strftime_now = time.strftime("%Y-%m-%d-%H-%M-%S", time.localtime())
|
||||
if not os.path.exists("figures"):
|
||||
os.mkdir("figures")
|
||||
env.save_fig("./figures/{}{}.png".format(strftime_now,int(raw_score)))
|
||||
print("save_figure: figures/{}{}.png".format(strftime_now,int(raw_score)))
|
||||
try:
|
||||
print("start try")
|
||||
# cost is the routing estimation based on the MST algorithm
|
||||
hpwl, cost = comp_res(placedb, env.node_pos, env.ratio)
|
||||
print("hpwl = {:.2f}\tcost = {:.2f}".format(hpwl, cost))
|
||||
except:
|
||||
assert False
|
||||
|
||||
if args.is_test:
|
||||
print("save node_pos")
|
||||
hpwl, cost = comp_res(placedb, env.node_pos, env.ratio)
|
||||
print("hpwl = {:.2f}\tcost = {:.2f}".format(hpwl, cost))
|
||||
print("time = {}s".format(end-start))
|
||||
pl_file_path = "{}-{}-{}.pl".format(benchmark, int(hpwl), time.strftime("%Y-%m-%d-%H-%M-%S", time.localtime()) )
|
||||
save_placement(pl_file_path, env.node_pos, env.ratio)
|
||||
strftime_now = time.strftime("%Y-%m-%d-%H-%M-%S", time.localtime())
|
||||
pl_path = 'gg_place_new/{}-{}-{}-{}.pl'.format(benchmark, strftime_now, int(hpwl), int(cost))
|
||||
fwrite_pl = open(pl_path, 'w')
|
||||
for node_name in env.node_pos:
|
||||
if node_name == "V":
|
||||
continue
|
||||
x, y, size_x, size_y = env.node_pos[node_name]
|
||||
x = x * env.ratio + placedb.node_info[node_name]['x'] /2.0
|
||||
y = y * env.ratio + placedb.node_info[node_name]['y'] /2.0
|
||||
fwrite_pl.write("{}\t{:.4f}\t{:.4f}\n".format(node_name, x, y))
|
||||
fwrite_pl.close()
|
||||
strftime_now = time.strftime("%Y-%m-%d-%H-%M-%S", time.localtime())
|
||||
env.save_fig("./figures/{}-{}-{}-{}.png".format(benchmark, strftime_now, int(hpwl), int(cost)))
|
||||
|
||||
training_records.append(TrainingRecord(i_epoch, running_reward))
|
||||
if i_epoch % 1 ==0:
|
||||
print("Epoch {}, Moving average score is: {:.2f} ".format(i_epoch, running_reward))
|
||||
fwrite.write("{},{},{:.2f},{}\n".format(i_epoch, score, running_reward, agent.training_step))
|
||||
fwrite.flush()
|
||||
writer.add_scalar('reward', running_reward, i_epoch)
|
||||
if running_reward > -100:
|
||||
print("Solved! Moving average score is now {}!".format(running_reward))
|
||||
env.close()
|
||||
agent.save_param()
|
||||
break
|
||||
if i_epoch % 100 == 0:
|
||||
if placed_num_macro is None:
|
||||
env.write_gl_file("./gl/{}{}.gl".format(strftime, int(score)))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
0
maskplace/ariane/__init__.py
Normal file
0
maskplace/ariane/__init__.py
Normal file
BIN
maskplace/ariane/__pycache__/__init__.cpython-37.pyc
Normal file
BIN
maskplace/ariane/__pycache__/__init__.cpython-37.pyc
Normal file
Binary file not shown.
BIN
maskplace/ariane/__pycache__/__init__.cpython-39.pyc
Normal file
BIN
maskplace/ariane/__pycache__/__init__.cpython-39.pyc
Normal file
Binary file not shown.
BIN
maskplace/ariane/__pycache__/laiyao_pb2.cpython-37.pyc
Normal file
BIN
maskplace/ariane/__pycache__/laiyao_pb2.cpython-37.pyc
Normal file
Binary file not shown.
BIN
maskplace/ariane/__pycache__/laiyao_pb2.cpython-39.pyc
Normal file
BIN
maskplace/ariane/__pycache__/laiyao_pb2.cpython-39.pyc
Normal file
Binary file not shown.
BIN
maskplace/ariane/__pycache__/read_info.cpython-37.pyc
Normal file
BIN
maskplace/ariane/__pycache__/read_info.cpython-37.pyc
Normal file
Binary file not shown.
BIN
maskplace/ariane/__pycache__/read_info.cpython-39.pyc
Normal file
BIN
maskplace/ariane/__pycache__/read_info.cpython-39.pyc
Normal file
Binary file not shown.
56
maskplace/ariane/graph.proto
Normal file
56
maskplace/ariane/graph.proto
Normal file
@ -0,0 +1,56 @@
|
||||
syntax = "proto3";
|
||||
|
||||
package tensorflow;
|
||||
|
||||
// import "tensorflow/core/framework/function.proto";
|
||||
// import "tensorflow/core/framework/node_def.proto";
|
||||
// import "tensorflow/core/framework/versions.proto";
|
||||
|
||||
option cc_enable_arenas = true;
|
||||
option java_outer_classname = "GraphProtos";
|
||||
option java_multiple_files = true;
|
||||
option java_package = "org.tensorflow.framework";
|
||||
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/graph_go_proto";
|
||||
|
||||
// Represents the graph of operations
|
||||
message GraphDef {
|
||||
repeated NodeDef node = 1;
|
||||
|
||||
// Compatibility versions of the graph. See core/public/version.h for version
|
||||
// history. The GraphDef version is distinct from the TensorFlow version, and
|
||||
// each release of TensorFlow will support a range of GraphDef versions.
|
||||
VersionDef versions = 4;
|
||||
|
||||
// Deprecated single version field; use versions above instead. Since all
|
||||
// GraphDef changes before "versions" was introduced were forward
|
||||
// compatible, this field is entirely ignored.
|
||||
int32 version = 3 [deprecated = true];
|
||||
|
||||
// "library" provides user-defined functions.
|
||||
//
|
||||
// Naming:
|
||||
// * library.function.name are in a flat namespace.
|
||||
// NOTE: We may need to change it to be hierarchical to support
|
||||
// different orgs. E.g.,
|
||||
// { "/google/nn", { ... }},
|
||||
// { "/google/vision", { ... }}
|
||||
// { "/org_foo/module_bar", { ... }}
|
||||
// map<string, FunctionDefLib> named_lib;
|
||||
// * If node[i].op is the name of one function in "library",
|
||||
// node[i] is deemed as a function call. Otherwise, node[i].op
|
||||
// must be a primitive operation supported by the runtime.
|
||||
//
|
||||
//
|
||||
// Function call semantics:
|
||||
//
|
||||
// * The callee may start execution as soon as some of its inputs
|
||||
// are ready. The caller may want to use Tuple() mechanism to
|
||||
// ensure all inputs are ready in the same time.
|
||||
//
|
||||
// * The consumer of return values may start executing as soon as
|
||||
// the return values the consumer depends on are ready. The
|
||||
// consumer may want to use Tuple() mechanism to ensure the
|
||||
// consumer does not start until all return values of the callee
|
||||
// function are ready.
|
||||
FunctionDefLibrary library = 2;
|
||||
}
|
||||
57
maskplace/ariane/laiyao.proto
Normal file
57
maskplace/ariane/laiyao.proto
Normal file
@ -0,0 +1,57 @@
|
||||
syntax = "proto3";
|
||||
|
||||
message AttrValue {
|
||||
// LINT.IfChange
|
||||
message ListValue {
|
||||
repeated bytes s = 2; // "list(string)"
|
||||
repeated int64 i = 3 [packed = true]; // "list(int)"
|
||||
repeated float f = 4 [packed = true]; // "list(float)"
|
||||
repeated bool b = 5 [packed = true]; // "list(bool)"
|
||||
// repeated DataType type = 6 [packed = true]; // "list(type)"
|
||||
// repeated TensorShapeProto shape = 7; // "list(shape)"
|
||||
// repeated TensorProto tensor = 8; // "list(tensor)"
|
||||
repeated NameAttrList func = 9; // "list(attr)"
|
||||
}
|
||||
// LINT.ThenChange(https://www.tensorflow.org/code/tensorflow/c/c_api.cc)
|
||||
|
||||
oneof value {
|
||||
bytes s = 2; // "string"
|
||||
int64 i = 3; // "int"
|
||||
float f = 4; // "float"
|
||||
bool b = 5; // "bool"
|
||||
DataType type = 6; // "type"
|
||||
// TensorShapeProto shape = 7; // "shape"
|
||||
// TensorProto tensor = 8; // "tensor"
|
||||
// ListValue list = 1; // any "list(...)"
|
||||
|
||||
// "func" represents a function. func.name is a function's name or
|
||||
// a primitive op's name. func.attr.first is the name of an attr
|
||||
// defined for that function. func.attr.second is the value for
|
||||
// that attr in the instantiation.
|
||||
NameAttrList func = 10;
|
||||
|
||||
// This is a placeholder only used in nodes defined inside a
|
||||
// function. It indicates the attr value will be supplied when
|
||||
// the function is instantiated. For example, let us suppose a
|
||||
// node "N" in function "FN". "N" has an attr "A" with value
|
||||
// placeholder = "foo". When FN is instantiated with attr "foo"
|
||||
// set to "bar", the instantiated node N's attr A will have been
|
||||
// given the value "bar".
|
||||
string placeholder = 9;
|
||||
}
|
||||
}
|
||||
|
||||
// A list of attr names and their values. The whole list is attached
|
||||
// with a string name. E.g., MatMul[T=float].
|
||||
message NameAttrList {
|
||||
string name = 1;
|
||||
map<string, AttrValue> attr = 2;
|
||||
}
|
||||
message NodeDef {
|
||||
string name = 1;
|
||||
repeated string input = 2;
|
||||
map<string, AttrValue> attr = 5;
|
||||
}
|
||||
message GraphDef {
|
||||
repeated NodeDef node = 1;
|
||||
}
|
||||
434
maskplace/ariane/laiyao_pb2.py
Normal file
434
maskplace/ariane/laiyao_pb2.py
Normal file
@ -0,0 +1,434 @@
|
||||
# Generated by the protocol buffer compiler. DO NOT EDIT!
|
||||
# source: laiyao.proto
|
||||
|
||||
import sys
|
||||
_b=sys.version_info[0]<3 and (lambda x:x) or (lambda x:x.encode('latin1'))
|
||||
from google.protobuf import descriptor as _descriptor
|
||||
from google.protobuf import message as _message
|
||||
from google.protobuf import reflection as _reflection
|
||||
from google.protobuf import symbol_database as _symbol_database
|
||||
# @@protoc_insertion_point(imports)
|
||||
|
||||
_sym_db = _symbol_database.Default()
|
||||
|
||||
|
||||
|
||||
|
||||
DESCRIPTOR = _descriptor.FileDescriptor(
|
||||
name='laiyao.proto',
|
||||
package='',
|
||||
syntax='proto3',
|
||||
serialized_options=None,
|
||||
serialized_pb=_b('\n\x0claiyao.proto\"\xe0\x01\n\tAttrValue\x12\x0b\n\x01s\x18\x02 \x01(\x0cH\x00\x12\x0b\n\x01i\x18\x03 \x01(\x03H\x00\x12\x0b\n\x01\x66\x18\x04 \x01(\x02H\x00\x12\x0b\n\x01\x62\x18\x05 \x01(\x08H\x00\x12\x1d\n\x04\x66unc\x18\n \x01(\x0b\x32\r.NameAttrListH\x00\x12\x15\n\x0bplaceholder\x18\t \x01(\tH\x00\x1a`\n\tListValue\x12\t\n\x01s\x18\x02 \x03(\x0c\x12\r\n\x01i\x18\x03 \x03(\x03\x42\x02\x10\x01\x12\r\n\x01\x66\x18\x04 \x03(\x02\x42\x02\x10\x01\x12\r\n\x01\x62\x18\x05 \x03(\x08\x42\x02\x10\x01\x12\x1b\n\x04\x66unc\x18\t \x03(\x0b\x32\r.NameAttrListB\x07\n\x05value\"|\n\x0cNameAttrList\x12\x0c\n\x04name\x18\x01 \x01(\t\x12%\n\x04\x61ttr\x18\x02 \x03(\x0b\x32\x17.NameAttrList.AttrEntry\x1a\x37\n\tAttrEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\x19\n\x05value\x18\x02 \x01(\x0b\x32\n.AttrValue:\x02\x38\x01\"\x81\x01\n\x07NodeDef\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\r\n\x05input\x18\x02 \x03(\t\x12 \n\x04\x61ttr\x18\x05 \x03(\x0b\x32\x12.NodeDef.AttrEntry\x1a\x37\n\tAttrEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\x19\n\x05value\x18\x02 \x01(\x0b\x32\n.AttrValue:\x02\x38\x01\"\"\n\x08GraphDef\x12\x16\n\x04node\x18\x01 \x03(\x0b\x32\x08.NodeDefb\x06proto3')
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
_ATTRVALUE_LISTVALUE = _descriptor.Descriptor(
|
||||
name='ListValue',
|
||||
full_name='AttrValue.ListValue',
|
||||
filename=None,
|
||||
file=DESCRIPTOR,
|
||||
containing_type=None,
|
||||
fields=[
|
||||
_descriptor.FieldDescriptor(
|
||||
name='s', full_name='AttrValue.ListValue.s', index=0,
|
||||
number=2, type=12, cpp_type=9, label=3,
|
||||
has_default_value=False, default_value=[],
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='i', full_name='AttrValue.ListValue.i', index=1,
|
||||
number=3, type=3, cpp_type=2, label=3,
|
||||
has_default_value=False, default_value=[],
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=_b('\020\001'), file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='f', full_name='AttrValue.ListValue.f', index=2,
|
||||
number=4, type=2, cpp_type=6, label=3,
|
||||
has_default_value=False, default_value=[],
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=_b('\020\001'), file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='b', full_name='AttrValue.ListValue.b', index=3,
|
||||
number=5, type=8, cpp_type=7, label=3,
|
||||
has_default_value=False, default_value=[],
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=_b('\020\001'), file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='func', full_name='AttrValue.ListValue.func', index=4,
|
||||
number=9, type=11, cpp_type=10, label=3,
|
||||
has_default_value=False, default_value=[],
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
],
|
||||
extensions=[
|
||||
],
|
||||
nested_types=[],
|
||||
enum_types=[
|
||||
],
|
||||
serialized_options=None,
|
||||
is_extendable=False,
|
||||
syntax='proto3',
|
||||
extension_ranges=[],
|
||||
oneofs=[
|
||||
],
|
||||
serialized_start=136,
|
||||
serialized_end=232,
|
||||
)
|
||||
|
||||
_ATTRVALUE = _descriptor.Descriptor(
|
||||
name='AttrValue',
|
||||
full_name='AttrValue',
|
||||
filename=None,
|
||||
file=DESCRIPTOR,
|
||||
containing_type=None,
|
||||
fields=[
|
||||
_descriptor.FieldDescriptor(
|
||||
name='s', full_name='AttrValue.s', index=0,
|
||||
number=2, type=12, cpp_type=9, label=1,
|
||||
has_default_value=False, default_value=_b(""),
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='i', full_name='AttrValue.i', index=1,
|
||||
number=3, type=3, cpp_type=2, label=1,
|
||||
has_default_value=False, default_value=0,
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='f', full_name='AttrValue.f', index=2,
|
||||
number=4, type=2, cpp_type=6, label=1,
|
||||
has_default_value=False, default_value=float(0),
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='b', full_name='AttrValue.b', index=3,
|
||||
number=5, type=8, cpp_type=7, label=1,
|
||||
has_default_value=False, default_value=False,
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='func', full_name='AttrValue.func', index=4,
|
||||
number=10, type=11, cpp_type=10, label=1,
|
||||
has_default_value=False, default_value=None,
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='placeholder', full_name='AttrValue.placeholder', index=5,
|
||||
number=9, type=9, cpp_type=9, label=1,
|
||||
has_default_value=False, default_value=_b("").decode('utf-8'),
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
],
|
||||
extensions=[
|
||||
],
|
||||
nested_types=[_ATTRVALUE_LISTVALUE, ],
|
||||
enum_types=[
|
||||
],
|
||||
serialized_options=None,
|
||||
is_extendable=False,
|
||||
syntax='proto3',
|
||||
extension_ranges=[],
|
||||
oneofs=[
|
||||
_descriptor.OneofDescriptor(
|
||||
name='value', full_name='AttrValue.value',
|
||||
index=0, containing_type=None, fields=[]),
|
||||
],
|
||||
serialized_start=17,
|
||||
serialized_end=241,
|
||||
)
|
||||
|
||||
|
||||
_NAMEATTRLIST_ATTRENTRY = _descriptor.Descriptor(
|
||||
name='AttrEntry',
|
||||
full_name='NameAttrList.AttrEntry',
|
||||
filename=None,
|
||||
file=DESCRIPTOR,
|
||||
containing_type=None,
|
||||
fields=[
|
||||
_descriptor.FieldDescriptor(
|
||||
name='key', full_name='NameAttrList.AttrEntry.key', index=0,
|
||||
number=1, type=9, cpp_type=9, label=1,
|
||||
has_default_value=False, default_value=_b("").decode('utf-8'),
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='value', full_name='NameAttrList.AttrEntry.value', index=1,
|
||||
number=2, type=11, cpp_type=10, label=1,
|
||||
has_default_value=False, default_value=None,
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
],
|
||||
extensions=[
|
||||
],
|
||||
nested_types=[],
|
||||
enum_types=[
|
||||
],
|
||||
serialized_options=_b('8\001'),
|
||||
is_extendable=False,
|
||||
syntax='proto3',
|
||||
extension_ranges=[],
|
||||
oneofs=[
|
||||
],
|
||||
serialized_start=312,
|
||||
serialized_end=367,
|
||||
)
|
||||
|
||||
_NAMEATTRLIST = _descriptor.Descriptor(
|
||||
name='NameAttrList',
|
||||
full_name='NameAttrList',
|
||||
filename=None,
|
||||
file=DESCRIPTOR,
|
||||
containing_type=None,
|
||||
fields=[
|
||||
_descriptor.FieldDescriptor(
|
||||
name='name', full_name='NameAttrList.name', index=0,
|
||||
number=1, type=9, cpp_type=9, label=1,
|
||||
has_default_value=False, default_value=_b("").decode('utf-8'),
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='attr', full_name='NameAttrList.attr', index=1,
|
||||
number=2, type=11, cpp_type=10, label=3,
|
||||
has_default_value=False, default_value=[],
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
],
|
||||
extensions=[
|
||||
],
|
||||
nested_types=[_NAMEATTRLIST_ATTRENTRY, ],
|
||||
enum_types=[
|
||||
],
|
||||
serialized_options=None,
|
||||
is_extendable=False,
|
||||
syntax='proto3',
|
||||
extension_ranges=[],
|
||||
oneofs=[
|
||||
],
|
||||
serialized_start=243,
|
||||
serialized_end=367,
|
||||
)
|
||||
|
||||
|
||||
_NODEDEF_ATTRENTRY = _descriptor.Descriptor(
|
||||
name='AttrEntry',
|
||||
full_name='NodeDef.AttrEntry',
|
||||
filename=None,
|
||||
file=DESCRIPTOR,
|
||||
containing_type=None,
|
||||
fields=[
|
||||
_descriptor.FieldDescriptor(
|
||||
name='key', full_name='NodeDef.AttrEntry.key', index=0,
|
||||
number=1, type=9, cpp_type=9, label=1,
|
||||
has_default_value=False, default_value=_b("").decode('utf-8'),
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='value', full_name='NodeDef.AttrEntry.value', index=1,
|
||||
number=2, type=11, cpp_type=10, label=1,
|
||||
has_default_value=False, default_value=None,
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
],
|
||||
extensions=[
|
||||
],
|
||||
nested_types=[],
|
||||
enum_types=[
|
||||
],
|
||||
serialized_options=_b('8\001'),
|
||||
is_extendable=False,
|
||||
syntax='proto3',
|
||||
extension_ranges=[],
|
||||
oneofs=[
|
||||
],
|
||||
serialized_start=312,
|
||||
serialized_end=367,
|
||||
)
|
||||
|
||||
_NODEDEF = _descriptor.Descriptor(
|
||||
name='NodeDef',
|
||||
full_name='NodeDef',
|
||||
filename=None,
|
||||
file=DESCRIPTOR,
|
||||
containing_type=None,
|
||||
fields=[
|
||||
_descriptor.FieldDescriptor(
|
||||
name='name', full_name='NodeDef.name', index=0,
|
||||
number=1, type=9, cpp_type=9, label=1,
|
||||
has_default_value=False, default_value=_b("").decode('utf-8'),
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='input', full_name='NodeDef.input', index=1,
|
||||
number=2, type=9, cpp_type=9, label=3,
|
||||
has_default_value=False, default_value=[],
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
_descriptor.FieldDescriptor(
|
||||
name='attr', full_name='NodeDef.attr', index=2,
|
||||
number=5, type=11, cpp_type=10, label=3,
|
||||
has_default_value=False, default_value=[],
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
],
|
||||
extensions=[
|
||||
],
|
||||
nested_types=[_NODEDEF_ATTRENTRY, ],
|
||||
enum_types=[
|
||||
],
|
||||
serialized_options=None,
|
||||
is_extendable=False,
|
||||
syntax='proto3',
|
||||
extension_ranges=[],
|
||||
oneofs=[
|
||||
],
|
||||
serialized_start=370,
|
||||
serialized_end=499,
|
||||
)
|
||||
|
||||
|
||||
_GRAPHDEF = _descriptor.Descriptor(
|
||||
name='GraphDef',
|
||||
full_name='GraphDef',
|
||||
filename=None,
|
||||
file=DESCRIPTOR,
|
||||
containing_type=None,
|
||||
fields=[
|
||||
_descriptor.FieldDescriptor(
|
||||
name='node', full_name='GraphDef.node', index=0,
|
||||
number=1, type=11, cpp_type=10, label=3,
|
||||
has_default_value=False, default_value=[],
|
||||
message_type=None, enum_type=None, containing_type=None,
|
||||
is_extension=False, extension_scope=None,
|
||||
serialized_options=None, file=DESCRIPTOR),
|
||||
],
|
||||
extensions=[
|
||||
],
|
||||
nested_types=[],
|
||||
enum_types=[
|
||||
],
|
||||
serialized_options=None,
|
||||
is_extendable=False,
|
||||
syntax='proto3',
|
||||
extension_ranges=[],
|
||||
oneofs=[
|
||||
],
|
||||
serialized_start=501,
|
||||
serialized_end=535,
|
||||
)
|
||||
|
||||
_ATTRVALUE_LISTVALUE.fields_by_name['func'].message_type = _NAMEATTRLIST
|
||||
_ATTRVALUE_LISTVALUE.containing_type = _ATTRVALUE
|
||||
_ATTRVALUE.fields_by_name['func'].message_type = _NAMEATTRLIST
|
||||
_ATTRVALUE.oneofs_by_name['value'].fields.append(
|
||||
_ATTRVALUE.fields_by_name['s'])
|
||||
_ATTRVALUE.fields_by_name['s'].containing_oneof = _ATTRVALUE.oneofs_by_name['value']
|
||||
_ATTRVALUE.oneofs_by_name['value'].fields.append(
|
||||
_ATTRVALUE.fields_by_name['i'])
|
||||
_ATTRVALUE.fields_by_name['i'].containing_oneof = _ATTRVALUE.oneofs_by_name['value']
|
||||
_ATTRVALUE.oneofs_by_name['value'].fields.append(
|
||||
_ATTRVALUE.fields_by_name['f'])
|
||||
_ATTRVALUE.fields_by_name['f'].containing_oneof = _ATTRVALUE.oneofs_by_name['value']
|
||||
_ATTRVALUE.oneofs_by_name['value'].fields.append(
|
||||
_ATTRVALUE.fields_by_name['b'])
|
||||
_ATTRVALUE.fields_by_name['b'].containing_oneof = _ATTRVALUE.oneofs_by_name['value']
|
||||
_ATTRVALUE.oneofs_by_name['value'].fields.append(
|
||||
_ATTRVALUE.fields_by_name['func'])
|
||||
_ATTRVALUE.fields_by_name['func'].containing_oneof = _ATTRVALUE.oneofs_by_name['value']
|
||||
_ATTRVALUE.oneofs_by_name['value'].fields.append(
|
||||
_ATTRVALUE.fields_by_name['placeholder'])
|
||||
_ATTRVALUE.fields_by_name['placeholder'].containing_oneof = _ATTRVALUE.oneofs_by_name['value']
|
||||
_NAMEATTRLIST_ATTRENTRY.fields_by_name['value'].message_type = _ATTRVALUE
|
||||
_NAMEATTRLIST_ATTRENTRY.containing_type = _NAMEATTRLIST
|
||||
_NAMEATTRLIST.fields_by_name['attr'].message_type = _NAMEATTRLIST_ATTRENTRY
|
||||
_NODEDEF_ATTRENTRY.fields_by_name['value'].message_type = _ATTRVALUE
|
||||
_NODEDEF_ATTRENTRY.containing_type = _NODEDEF
|
||||
_NODEDEF.fields_by_name['attr'].message_type = _NODEDEF_ATTRENTRY
|
||||
_GRAPHDEF.fields_by_name['node'].message_type = _NODEDEF
|
||||
DESCRIPTOR.message_types_by_name['AttrValue'] = _ATTRVALUE
|
||||
DESCRIPTOR.message_types_by_name['NameAttrList'] = _NAMEATTRLIST
|
||||
DESCRIPTOR.message_types_by_name['NodeDef'] = _NODEDEF
|
||||
DESCRIPTOR.message_types_by_name['GraphDef'] = _GRAPHDEF
|
||||
_sym_db.RegisterFileDescriptor(DESCRIPTOR)
|
||||
|
||||
AttrValue = _reflection.GeneratedProtocolMessageType('AttrValue', (_message.Message,), dict(
|
||||
|
||||
ListValue = _reflection.GeneratedProtocolMessageType('ListValue', (_message.Message,), dict(
|
||||
DESCRIPTOR = _ATTRVALUE_LISTVALUE,
|
||||
__module__ = 'laiyao_pb2'
|
||||
# @@protoc_insertion_point(class_scope:AttrValue.ListValue)
|
||||
))
|
||||
,
|
||||
DESCRIPTOR = _ATTRVALUE,
|
||||
__module__ = 'laiyao_pb2'
|
||||
# @@protoc_insertion_point(class_scope:AttrValue)
|
||||
))
|
||||
_sym_db.RegisterMessage(AttrValue)
|
||||
_sym_db.RegisterMessage(AttrValue.ListValue)
|
||||
|
||||
NameAttrList = _reflection.GeneratedProtocolMessageType('NameAttrList', (_message.Message,), dict(
|
||||
|
||||
AttrEntry = _reflection.GeneratedProtocolMessageType('AttrEntry', (_message.Message,), dict(
|
||||
DESCRIPTOR = _NAMEATTRLIST_ATTRENTRY,
|
||||
__module__ = 'laiyao_pb2'
|
||||
# @@protoc_insertion_point(class_scope:NameAttrList.AttrEntry)
|
||||
))
|
||||
,
|
||||
DESCRIPTOR = _NAMEATTRLIST,
|
||||
__module__ = 'laiyao_pb2'
|
||||
# @@protoc_insertion_point(class_scope:NameAttrList)
|
||||
))
|
||||
_sym_db.RegisterMessage(NameAttrList)
|
||||
_sym_db.RegisterMessage(NameAttrList.AttrEntry)
|
||||
|
||||
NodeDef = _reflection.GeneratedProtocolMessageType('NodeDef', (_message.Message,), dict(
|
||||
|
||||
AttrEntry = _reflection.GeneratedProtocolMessageType('AttrEntry', (_message.Message,), dict(
|
||||
DESCRIPTOR = _NODEDEF_ATTRENTRY,
|
||||
__module__ = 'laiyao_pb2'
|
||||
# @@protoc_insertion_point(class_scope:NodeDef.AttrEntry)
|
||||
))
|
||||
,
|
||||
DESCRIPTOR = _NODEDEF,
|
||||
__module__ = 'laiyao_pb2'
|
||||
# @@protoc_insertion_point(class_scope:NodeDef)
|
||||
))
|
||||
_sym_db.RegisterMessage(NodeDef)
|
||||
_sym_db.RegisterMessage(NodeDef.AttrEntry)
|
||||
|
||||
GraphDef = _reflection.GeneratedProtocolMessageType('GraphDef', (_message.Message,), dict(
|
||||
DESCRIPTOR = _GRAPHDEF,
|
||||
__module__ = 'laiyao_pb2'
|
||||
# @@protoc_insertion_point(class_scope:GraphDef)
|
||||
))
|
||||
_sym_db.RegisterMessage(GraphDef)
|
||||
|
||||
|
||||
_ATTRVALUE_LISTVALUE.fields_by_name['i']._options = None
|
||||
_ATTRVALUE_LISTVALUE.fields_by_name['f']._options = None
|
||||
_ATTRVALUE_LISTVALUE.fields_by_name['b']._options = None
|
||||
_NAMEATTRLIST_ATTRENTRY._options = None
|
||||
_NODEDEF_ATTRENTRY._options = None
|
||||
# @@protoc_insertion_point(module_scope)
|
||||
997324
maskplace/ariane/netlist.pb.txt
Normal file
997324
maskplace/ariane/netlist.pb.txt
Normal file
File diff suppressed because it is too large
Load Diff
62
maskplace/ariane/read_info.py
Normal file
62
maskplace/ariane/read_info.py
Normal file
@ -0,0 +1,62 @@
|
||||
# import tensorflow as tf
|
||||
|
||||
from google.protobuf import text_format
|
||||
|
||||
import laiyao_pb2
|
||||
|
||||
|
||||
def load_pbtxt_file(path):
|
||||
"""Read .pbtxt file.
|
||||
|
||||
Args:
|
||||
path: Path to StringIntLabelMap proto text file (.pbtxt file).
|
||||
|
||||
Returns:
|
||||
A StringIntLabelMapProto.
|
||||
|
||||
Raises:
|
||||
ValueError: If path is not exist.
|
||||
"""
|
||||
# if not tf.gfile.Exists(path):
|
||||
# raise ValueError('`path` is not exist.')
|
||||
|
||||
# with tf.gfile.GFile(path, 'r') as fid:
|
||||
# pbtxt_string = fid.read()
|
||||
# pbtxt = laiyao_pb2.StudentInfo()
|
||||
# try:
|
||||
# text_format.Merge(pbtxt_string, pbtxt)
|
||||
# except text_format.ParseError:
|
||||
# pbtxt.ParseFromString(pbtxt_string)
|
||||
fid = open(path, 'r')
|
||||
pbtxt_string = fid.read()
|
||||
pbtxt = laiyao_pb2.GraphDef()
|
||||
try:
|
||||
text_format.Merge(pbtxt_string, pbtxt)
|
||||
except text_format.ParseError:
|
||||
pbtxt.ParseFromString(pbtxt_string)
|
||||
return pbtxt
|
||||
|
||||
|
||||
def get_netlist_info_dict(path):
|
||||
"""Reads a .pbtxt file and returns a dictionary.
|
||||
|
||||
Args:
|
||||
path: Path to StringIntLabelMap proto text file.
|
||||
|
||||
Returns:
|
||||
A dictionary mapping class names to indices.
|
||||
"""
|
||||
pbtxt = load_pbtxt_file(path)
|
||||
|
||||
# result_dict = {}
|
||||
# for node in pbtxt.node:
|
||||
# print("node_name: {}".format(node.name))
|
||||
return pbtxt
|
||||
|
||||
|
||||
def main():
|
||||
get_netlist_info_dict('netlist.pb.txt')
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
51
maskplace/comp_res.py
Normal file
51
maskplace/comp_res.py
Normal file
@ -0,0 +1,51 @@
|
||||
from place_db import PlaceDB
|
||||
from prim import prim_real
|
||||
import pickle
|
||||
|
||||
def comp_res(placedb, node_pos, ratio):
|
||||
|
||||
hpwl = 0.0
|
||||
cost = 0.0
|
||||
for net_name in placedb.net_info:
|
||||
max_x = 0.0
|
||||
min_x = placedb.max_height * 1.1
|
||||
max_y = 0.0
|
||||
min_y = placedb.max_height * 1.1
|
||||
for node_name in placedb.net_info[net_name]["nodes"]:
|
||||
if node_name not in node_pos:
|
||||
continue
|
||||
h = placedb.node_info[node_name]['x']
|
||||
w = placedb.node_info[node_name]['y']
|
||||
pin_x = node_pos[node_name][0] * ratio + h / 2.0 + placedb.net_info[net_name]["nodes"][node_name]["x_offset"]
|
||||
pin_y = node_pos[node_name][1] * ratio + w / 2.0 + placedb.net_info[net_name]["nodes"][node_name]["y_offset"]
|
||||
max_x = max(pin_x, max_x)
|
||||
min_x = min(pin_x, min_x)
|
||||
max_y = max(pin_y, max_y)
|
||||
min_y = min(pin_y, min_y)
|
||||
for port_name in placedb.net_info[net_name]["ports"]:
|
||||
h = placedb.port_info[port_name]['x']
|
||||
w = placedb.port_info[port_name]['y']
|
||||
pin_x = h
|
||||
pin_y = w
|
||||
max_x = max(pin_x, max_x)
|
||||
min_x = min(pin_x, min_x)
|
||||
max_y = max(pin_y, max_y)
|
||||
min_y = min(pin_y, min_y)
|
||||
if min_x <= placedb.max_height:
|
||||
hpwl_tmp = (max_x - min_x) + (max_y - min_y)
|
||||
else:
|
||||
hpwl_tmp = 0
|
||||
if "weight" in placedb.net_info[net_name]:
|
||||
hpwl_tmp *= placedb.net_info[net_name]["weight"]
|
||||
hpwl += hpwl_tmp
|
||||
net_node_set = set.union(set(placedb.net_info[net_name]["nodes"]),
|
||||
set(placedb.net_info[net_name]["ports"]))
|
||||
for net_node in list(net_node_set):
|
||||
if net_node not in node_pos and net_node not in placedb.port_info:
|
||||
net_node_set.discard(net_node)
|
||||
prim_cost = prim_real(net_node_set, node_pos, placedb.net_info[net_name]["nodes"], ratio, placedb.node_info, placedb.port_info)
|
||||
if "weight" in placedb.net_info[net_name]:
|
||||
prim_cost *= placedb.net_info[net_name]["weight"]
|
||||
assert hpwl_tmp <= prim_cost +1e-5
|
||||
cost += prim_cost
|
||||
return hpwl, cost
|
||||
@ -1,23 +1,22 @@
|
||||
import numpy as np
|
||||
import os
|
||||
import random
|
||||
from operator import itemgetter
|
||||
from itertools import combinations
|
||||
|
||||
from place_db_proto import get_node_info
|
||||
from place_db_proto import get_net_info
|
||||
import sys
|
||||
if os.path.exists('benchmark/ariane'):
|
||||
sys.path.append('benchmark/ariane')
|
||||
from ariane.read_info import get_netlist_info_dict
|
||||
from place_db_proto import get_node_info
|
||||
from place_db_proto import get_net_info
|
||||
|
||||
import pickle
|
||||
sys.path.append('ariane')
|
||||
from ariane.read_info import get_netlist_info_dict
|
||||
# Macro dict (macro id -> name, x, y)
|
||||
|
||||
def read_node_file(fopen, benchmark):
|
||||
node_info = {}
|
||||
node_info_raw_id_name ={}
|
||||
port_info = {}
|
||||
node_cnt = 0
|
||||
for line in fopen.readlines():
|
||||
if not line.startswith("\t") and not line.startswith(" "):
|
||||
if not line.startswith("\t"):
|
||||
continue
|
||||
line = line.strip().split()
|
||||
if line[-1] != "terminal":
|
||||
@ -28,9 +27,8 @@ def read_node_file(fopen, benchmark):
|
||||
node_info[node_name] = {"id": node_cnt, "x": x , "y": y }
|
||||
node_info_raw_id_name[node_cnt] = node_name
|
||||
node_cnt += 1
|
||||
|
||||
print("len node_info", len(node_info))
|
||||
return node_info, node_info_raw_id_name, port_info
|
||||
return node_info, node_info_raw_id_name
|
||||
|
||||
|
||||
def read_net_file(fopen, node_info):
|
||||
@ -38,8 +36,7 @@ def read_net_file(fopen, node_info):
|
||||
net_name = None
|
||||
net_cnt = 0
|
||||
for line in fopen.readlines():
|
||||
if not line.startswith("\t") and not line.startswith(" ") and \
|
||||
not line.startswith("NetDegree"):
|
||||
if not line.startswith("\t") and not line.startswith("NetDegree"):
|
||||
continue
|
||||
line = line.strip().split()
|
||||
if line[0] == "NetDegree":
|
||||
@ -51,16 +48,11 @@ def read_net_file(fopen, node_info):
|
||||
net_info[net_name] = {}
|
||||
net_info[net_name]["nodes"] = {}
|
||||
net_info[net_name]["ports"] = {}
|
||||
if not node_name.startswith("p") and not node_name in net_info[net_name]["nodes"]:
|
||||
if not node_name in net_info[net_name]["nodes"]:
|
||||
x_offset = float(line[-2])
|
||||
y_offset = float(line[-1])
|
||||
net_info[net_name]["nodes"][node_name] = {}
|
||||
net_info[net_name]["nodes"][node_name] = {"x_offset": x_offset, "y_offset": y_offset}
|
||||
elif node_name.startswith("p") and node_name in net_info[net_name]["ports"]:
|
||||
x_offset = float(line[-2])
|
||||
y_offset = float(line[-1])
|
||||
net_info[net_name]["ports"][node_name] = {}
|
||||
net_info[net_name]["ports"][node_name] = {"x_offset": x_offset, "y_offset": y_offset}
|
||||
for net_name in list(net_info.keys()):
|
||||
if len(net_info[net_name]["nodes"]) <= 1:
|
||||
net_info.pop(net_name)
|
||||
@ -102,7 +94,6 @@ def get_port_to_net_dict(port_info, net_info):
|
||||
port_to_net_dict[port_name].add(net_name)
|
||||
return port_to_net_dict
|
||||
|
||||
|
||||
def read_pl_file(fopen, node_info):
|
||||
max_height = 0
|
||||
max_width = 0
|
||||
@ -122,17 +113,6 @@ def read_pl_file(fopen, node_info):
|
||||
return max(max_height, max_width), max(max_height, max_width)
|
||||
|
||||
|
||||
def read_scl_file(fopen, benchmark):
|
||||
assert "ibm" in benchmark
|
||||
for line in fopen.readlines():
|
||||
if not "Numsites" in line:
|
||||
continue
|
||||
line = line.strip().split()
|
||||
max_height = int(line[-1])
|
||||
break
|
||||
return max_height, max_height
|
||||
|
||||
|
||||
def get_node_id_to_name(node_info, node_to_net_dict):
|
||||
node_name_and_num = []
|
||||
for node_name in node_info:
|
||||
@ -203,43 +183,40 @@ def get_node_id_to_name_topology(node_info, node_to_net_dict, net_info, benchmar
|
||||
candidates[node_name] = 0
|
||||
if len(candidates) > 0:
|
||||
if benchmark != 'ariane':
|
||||
add_node = max(candidates, key = lambda v: candidates[v]*1 + node_net_num[v]*1000 +\
|
||||
node_info[v]['x']*node_info[v]['y'] * 1 +int(hash(v)%10000)*1e-6)
|
||||
if benchmark == "bigblue3":
|
||||
add_node = max(candidates, key = lambda v: candidates[v]*1 + node_net_num[v]*100000 +\
|
||||
node_info[v]['x']*node_info[v]['y'] * 1 +int(hash(v)%10000)*1e-6)
|
||||
else:
|
||||
add_node = max(candidates, key = lambda v: candidates[v]*1 + node_net_num[v]*1000 +\
|
||||
node_info[v]['x']*node_info[v]['y'] * 1 +int(hash(v)%10000)*1e-6)
|
||||
else:
|
||||
add_node = max(candidates, key = lambda v: candidates[v]*30000 + node_net_num[v]*1000 +\
|
||||
node_info[v]['x']*node_info[v]['y']*1 +int(hash(v)%10000)*1e-6)
|
||||
else:
|
||||
add_node = max(node_net_num, key = lambda v: node_net_num[v]*1000 + node_info[v]['x']*node_info[v]['y']*1)
|
||||
if benchmark != 'ariane':
|
||||
if benchmark == "bigblue3":
|
||||
add_node = max(node_net_num, key = lambda v: node_net_num[v]*100000 + node_info[v]['x']*node_info[v]['y']*1)
|
||||
else:
|
||||
add_node = max(node_net_num, key = lambda v: node_net_num[v]*1000 + node_info[v]['x']*node_info[v]['y']*1)
|
||||
else:
|
||||
add_node = max(node_net_num, key = lambda v: node_net_num[v]*1000 + node_info[v]['x']*node_info[v]['y']*1)
|
||||
|
||||
visited_node.add(add_node)
|
||||
node_id_to_name.append((add_node, node_net_num[add_node]))
|
||||
node_id_to_name.append((add_node, node_net_num[add_node]))
|
||||
node_net_num.pop(add_node)
|
||||
for i, (node_name, _) in enumerate(node_id_to_name):
|
||||
node_info[node_name]["id"] = i
|
||||
print("node_id_to_name")
|
||||
print(node_id_to_name)
|
||||
# print("node_id_to_name")
|
||||
# print(node_id_to_name)
|
||||
node_id_to_name_res = [x for x, _ in node_id_to_name]
|
||||
return node_id_to_name_res
|
||||
|
||||
|
||||
def get_pin_cnt(net_info):
|
||||
pin_cnt = 0
|
||||
for net_name in net_info:
|
||||
pin_cnt += len(net_info[net_name]["nodes"])
|
||||
return pin_cnt
|
||||
|
||||
|
||||
def get_total_area(node_info):
|
||||
area = 0
|
||||
for node_name in node_info:
|
||||
area += node_info[node_name]["x"] * node_info[node_name]["y"]
|
||||
return area
|
||||
|
||||
|
||||
class PlaceDB():
|
||||
|
||||
def __init__(self, benchmark = "adaptec1"):
|
||||
self.benchmark = benchmark
|
||||
if benchmark == "ariane":
|
||||
if benchmark == "ariane" or benchmark == "sample_clustered":
|
||||
path = benchmark + '/netlist.pb.txt'
|
||||
pbtxt = get_netlist_info_dict(path)
|
||||
self.node_info, self.node_info_raw_id_name = get_node_info(pbtxt)
|
||||
@ -249,41 +226,32 @@ class PlaceDB():
|
||||
self.max_height, self.max_width = 357, 357
|
||||
self.port_to_net_dict = get_port_to_net_dict(self.port_info, self.net_info)
|
||||
else:
|
||||
assert os.path.exists(os.path.join("benchmark", benchmark))
|
||||
node_file = open(os.path.join("benchmark", benchmark, benchmark+".nodes"), "r")
|
||||
self.node_info, self.node_info_raw_id_name, self.port_info = read_node_file(node_file, benchmark)
|
||||
pl_file = open(os.path.join("benchmark", benchmark, benchmark+".pl"), "r")
|
||||
assert os.path.exists(benchmark)
|
||||
node_file = open(os.path.join(benchmark, benchmark+".nodes"), "r")
|
||||
self.node_info, self.node_info_raw_id_name = read_node_file(node_file, benchmark)
|
||||
pl_file = open(os.path.join(benchmark, benchmark+".pl"), "r")
|
||||
self.port_info = {}
|
||||
self.node_cnt = len(self.node_info)
|
||||
node_file.close()
|
||||
net_file = open(os.path.join("benchmark", benchmark, benchmark+".nets"), "r")
|
||||
net_file = open(os.path.join(benchmark, benchmark+".nets"), "r")
|
||||
self.net_info = read_net_file(net_file, self.node_info)
|
||||
self.net_cnt = len(self.net_info)
|
||||
net_file.close()
|
||||
pl_file = open(os.path.join("benchmark", benchmark, benchmark+".pl"), "r")
|
||||
pl_file = open(os.path.join(benchmark, benchmark+".pl"), "r")
|
||||
self.max_height, self.max_width = read_pl_file(pl_file, self.node_info)
|
||||
pl_file.close()
|
||||
if not "ibm" in benchmark:
|
||||
self.port_to_net_dict = {}
|
||||
else:
|
||||
self.port_to_net_dict = get_port_to_net_dict(self.port_info, self.net_info)
|
||||
scl_file = open(os.path.join("benchmark", benchmark, benchmark+".scl"), "r")
|
||||
self.max_height, self.max_width = read_scl_file(scl_file, benchmark)
|
||||
|
||||
self.port_to_net_dict = {}
|
||||
self.node_to_net_dict = get_node_to_net_dict(self.node_info, self.net_info)
|
||||
self.node_id_to_name = get_node_id_to_name_topology(self.node_info, self.node_to_net_dict, self.net_info, self.benchmark)
|
||||
self.node_name_to_id = dict((t, i) for i, t in enumerate(self.node_id_to_name))
|
||||
|
||||
|
||||
def debug_str(self):
|
||||
print("node_cnt = {}".format(len(self.node_info)))
|
||||
print("net_cnt = {}".format(len(self.net_info)))
|
||||
print("max_height = {}".format(self.max_height))
|
||||
print("max_width = {}".format(self.max_width))
|
||||
print("pin_cnt = {}".format(get_pin_cnt(self.net_info)))
|
||||
print("port_cnt = {}".format(len(self.port_info)))
|
||||
print("area_ratio = {}".format(get_total_area(self.node_info)/(self.max_height*self.max_height)))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
placedb = PlaceDB("adaptec1")
|
||||
placedb = PlaceDB("ariane")
|
||||
placedb.debug_str()
|
||||
|
||||
|
||||
117
maskplace/place_db_proto.py
Normal file
117
maskplace/place_db_proto.py
Normal file
@ -0,0 +1,117 @@
|
||||
import sys
|
||||
sys.path.append('ariane')
|
||||
from ariane.read_info import get_netlist_info_dict
|
||||
from tqdm import tqdm
|
||||
|
||||
def get_node_info(pbtxt):
|
||||
node_info = {}
|
||||
node_info_raw_id_name = {}
|
||||
node_cnt = 0
|
||||
area_sum = 0.0
|
||||
|
||||
for node in pbtxt.node:
|
||||
if node.attr['type'].placeholder.upper() != "MACRO":
|
||||
continue
|
||||
node_name = node.name
|
||||
x = float(node.attr['width'].f)
|
||||
y = float(node.attr['height'].f)
|
||||
node_info[node_name] = {"id": node_cnt, "x": x, "y": y}
|
||||
area_sum += x * y
|
||||
if node.attr['type'].placeholder == "MACRO":
|
||||
node_info[node_name]["is_hard"] = 1
|
||||
else:
|
||||
node_info[node_name]["is_hard"] = 0
|
||||
node_info_raw_id_name[node_cnt] = node_name
|
||||
node_cnt += 1
|
||||
print("area_sum = {}".format(area_sum))
|
||||
return node_info, node_info_raw_id_name
|
||||
|
||||
|
||||
def get_net_info(pbtxt):
|
||||
net_info = {}
|
||||
net_name = None
|
||||
net_cnt = 0
|
||||
pin_cnt = 0
|
||||
pin_info = {}
|
||||
port_info = {}
|
||||
for node in pbtxt.node:
|
||||
if node.attr['type'].placeholder.upper() == "MACRO":
|
||||
continue
|
||||
pin_name = node.name
|
||||
if node.attr['type'].placeholder.upper() == "PORT":
|
||||
x = float(node.attr['x'].f)
|
||||
y = float(node.attr['y'].f)
|
||||
port_info[pin_name] = {"x": x, "y": y}
|
||||
elif node.attr['type'].placeholder.upper() == "MACRO_PIN":
|
||||
macro_name = node.attr['macro_name'].placeholder
|
||||
x_offset = float(node.attr['x_offset'].f)
|
||||
y_offset = float(node.attr['y_offset'].f)
|
||||
pin_info[pin_name] = {"node_name": macro_name, "x_offset": x_offset, "y_offset": y_offset}
|
||||
pin_cnt += 1
|
||||
print("pin_cnt = {}".format(pin_cnt))
|
||||
for node in pbtxt.node:
|
||||
net_name = node.name
|
||||
if node.attr['type'].placeholder.upper() == "MACRO":
|
||||
continue
|
||||
net_info[net_name] = {}
|
||||
net_info[net_name]["nodes"] = {}
|
||||
net_info[net_name]["ports"] = {}
|
||||
if 'weight' in node.attr:
|
||||
net_info[net_name]["weight"] = float(node.attr['weight'].f)
|
||||
else:
|
||||
net_info[net_name]["weight"] = 1.0
|
||||
for pin_name in node.input:
|
||||
if pin_name in port_info:
|
||||
assert pin_name not in net_info[net_name]["ports"]
|
||||
net_info[net_name]["ports"][pin_name] = {}
|
||||
net_info[net_name]["ports"][pin_name]["x"] = port_info[pin_name]["x"]
|
||||
net_info[net_name]["ports"][pin_name]["y"] = port_info[pin_name]["y"]
|
||||
elif pin_name in pin_info:
|
||||
node_name = pin_info[pin_name]["node_name"]
|
||||
if node_name in net_info[net_name]["nodes"]:
|
||||
if "x_offsets" not in net_info[net_name]["nodes"][node_name]:
|
||||
net_info[net_name]["nodes"][node_name]["x_offsets"] = [net_info[net_name]["nodes"][node_name]["x_offset"]]
|
||||
net_info[net_name]["nodes"][node_name]["y_offsets"] = [net_info[net_name]["nodes"][node_name]["y_offset"]]
|
||||
net_info[net_name]["nodes"][node_name]["x_offsets"].append(pin_info[pin_name]["x_offset"])
|
||||
net_info[net_name]["nodes"][node_name]["y_offsets"].append(pin_info[pin_name]["y_offset"])
|
||||
net_info[net_name]["nodes"][node_name] = {}
|
||||
net_info[net_name]["nodes"][node_name]["x_offset"] = pin_info[pin_name]["x_offset"]
|
||||
net_info[net_name]["nodes"][node_name]["y_offset"] = pin_info[pin_name]["y_offset"]
|
||||
else:
|
||||
assert False
|
||||
out_pin_name = net_name
|
||||
if out_pin_name in port_info:
|
||||
assert out_pin_name not in net_info[net_name]["ports"]
|
||||
net_info[net_name]["ports"][out_pin_name] = {}
|
||||
net_info[net_name]["ports"][out_pin_name]["x"] = port_info[out_pin_name]["x"]
|
||||
net_info[net_name]["ports"][out_pin_name]["y"] = port_info[out_pin_name]["y"]
|
||||
elif out_pin_name in pin_info:
|
||||
node_name = pin_info[out_pin_name]["node_name"]
|
||||
assert node_name not in net_info[net_name]["nodes"]
|
||||
net_info[net_name]["nodes"][node_name] = {}
|
||||
net_info[net_name]["nodes"][node_name]["x_offset"] = pin_info[out_pin_name]["x_offset"]
|
||||
net_info[net_name]["nodes"][node_name]["y_offset"] = pin_info[out_pin_name]["y_offset"]
|
||||
else:
|
||||
print("out_pin_name = {}".format(out_pin_name))
|
||||
assert False
|
||||
|
||||
for net_name in list(net_info.keys()):
|
||||
if len(net_info[net_name]["nodes"]) + \
|
||||
len(net_info[net_name]["ports"]) <= 1:
|
||||
net_info.pop(net_name)
|
||||
for net_name in net_info:
|
||||
net_info[net_name]['id'] = net_cnt
|
||||
net_cnt += 1
|
||||
print("adjust net size = {}".format(len(net_info)))
|
||||
return net_info, port_info
|
||||
|
||||
|
||||
def main():
|
||||
path = 'ariane/netlist.pb.txt'
|
||||
pbtxt = get_netlist_info_dict(path)
|
||||
node_info = get_node_info(pbtxt)
|
||||
net_info, port_info = get_net_info(pbtxt)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
6
maskplace/place_env/__init__.py
Normal file
6
maskplace/place_env/__init__.py
Normal file
@ -0,0 +1,6 @@
|
||||
from gym.envs.registration import register
|
||||
|
||||
register(
|
||||
id = 'place_env-v0',
|
||||
entry_point = 'place_env.place_env:PlaceEnv'
|
||||
)
|
||||
301
maskplace/place_env/place_env.py
Normal file
301
maskplace/place_env/place_env.py
Normal file
@ -0,0 +1,301 @@
|
||||
import math
|
||||
import gym
|
||||
from gym import spaces
|
||||
import numpy as np
|
||||
import sys
|
||||
sys.path.append("..")
|
||||
from place_db import PlaceDB
|
||||
import matplotlib.pyplot as plt
|
||||
import matplotlib.patches as patches
|
||||
import time
|
||||
|
||||
class PlaceEnv(gym.Env):
|
||||
|
||||
def __init__(self, placedb, placed_num_macro = None, grid = 224):
|
||||
|
||||
# need to get GCN vector and CNN
|
||||
print("grid * grid", grid * grid)
|
||||
print("placedb.node_cnt", placedb.node_cnt)
|
||||
print("placedb.net_cnt", placedb.net_cnt)
|
||||
assert grid * grid >= placedb.node_cnt
|
||||
self.grid = grid
|
||||
self.max_height = placedb.max_height
|
||||
self.max_width = placedb.max_width
|
||||
self.placedb = placedb
|
||||
self.num_macro = placedb.node_cnt
|
||||
self.placed_num_macro = placed_num_macro
|
||||
self.num_net = placedb.net_cnt
|
||||
self.node_name_list = placedb.node_id_to_name
|
||||
self.action_space = spaces.Discrete(self.grid * self.grid)
|
||||
self.state = None
|
||||
self.net_min_max_ord = {}
|
||||
self.node_pos = {}
|
||||
self.net_placed_set = {}
|
||||
self.last_reward = 0
|
||||
self.num_macro_placed = 0
|
||||
self.node_x_max = 0
|
||||
self.node_x_min = self.grid
|
||||
self.node_y_max = 0
|
||||
self.node_y_min = self.grid
|
||||
self.ratio = self.placedb.max_height / self.grid
|
||||
print("self.ratio = {:.2f}".format(self.ratio))
|
||||
|
||||
def reset(self):
|
||||
self.num_macro_placed = 0
|
||||
num_macro = self.num_macro
|
||||
canvas = np.zeros((self.grid, self.grid))
|
||||
self.node_pos = {}
|
||||
self.net_min_max_ord = {}
|
||||
self.net_fea = np.zeros((self.num_net, 4))
|
||||
self.net_fea[:, 0] = 0
|
||||
self.net_fea[:, 1] = 1.0
|
||||
self.net_fea[:, 2] = 0
|
||||
self.net_fea[:, 3] = 1.0
|
||||
self.rudy = np.zeros((self.grid, self.grid))
|
||||
for port_name in self.placedb.port_to_net_dict:
|
||||
for net_name in self.placedb.port_to_net_dict[port_name]:
|
||||
pin_x = round(self.placedb.port_info[port_name]['x'] / self.ratio)
|
||||
pin_y = round(self.placedb.port_info[port_name]['y'] / self.ratio)
|
||||
if net_name in self.net_min_max_ord:
|
||||
if pin_x > self.net_min_max_ord[net_name]['max_x']:
|
||||
self.net_min_max_ord[net_name]['max_x'] = pin_x
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][1] = pin_x / self.grid
|
||||
elif pin_x < self.net_min_max_ord[net_name]['min_x']:
|
||||
self.net_min_max_ord[net_name]['max_y'] = pin_y
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][0] = pin_x / self.grid
|
||||
if pin_y > self.net_min_max_ord[net_name]['max_y']:
|
||||
self.net_min_max_ord[net_name]['max_y'] = pin_y
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][3] = pin_y / self.grid
|
||||
elif pin_y < self.net_min_max_ord[net_name]['min_y']:
|
||||
self.net_min_max_ord[net_name]['min_y'] = pin_y
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][2] = pin_y / self.grid
|
||||
else:
|
||||
self.net_min_max_ord[net_name] = {}
|
||||
self.net_min_max_ord[net_name]['max_x'] = pin_x
|
||||
self.net_min_max_ord[net_name]['min_x'] = pin_x
|
||||
self.net_min_max_ord[net_name]['max_y'] = pin_y
|
||||
self.net_min_max_ord[net_name]['min_y'] = pin_y
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][1] = pin_x / self.grid
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][0] = pin_x / self.grid
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][3] = pin_y / self.grid
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][2] = pin_y / self.grid
|
||||
self.net_placed_set = {}
|
||||
self.num_macro_placed = 0
|
||||
|
||||
net_img = np.zeros((self.grid, self.grid))
|
||||
net_img_2 = np.zeros((self.grid, self.grid))
|
||||
|
||||
next_x = math.ceil(max(1, self.placedb.node_info[self.node_name_list[self.num_macro_placed]]['x'] / self.ratio))
|
||||
next_y = math.ceil(max(1, self.placedb.node_info[self.node_name_list[self.num_macro_placed]]['y'] / self.ratio))
|
||||
mask = self.get_mask(canvas, next_x, next_y)
|
||||
next_x_2 = math.ceil(max(1, self.placedb.node_info[self.node_name_list[self.num_macro_placed+1]]['x'] / self.ratio))
|
||||
next_y_2 = math.ceil(max(1, self.placedb.node_info[self.node_name_list[self.num_macro_placed+1]]['y'] / self.ratio))
|
||||
mask_2 = self.get_mask(canvas, next_x_2, next_y_2)
|
||||
for net_name in self.placedb.net_info:
|
||||
self.net_placed_set[net_name] = set()
|
||||
self.state = np.concatenate((np.array([self.num_macro_placed]), canvas.flatten(),
|
||||
net_img.flatten(), mask.flatten(), net_img_2.flatten(), mask_2.flatten(),
|
||||
np.array([next_x/self.grid, next_y/self.grid])), axis = 0)
|
||||
self.node_x_max = 0
|
||||
self.node_x_min = self.grid
|
||||
self.node_y_max = 0
|
||||
self.node_y_min = self.grid
|
||||
|
||||
return self.state
|
||||
|
||||
def save_fig(self, file_path):
|
||||
fig1 = plt.figure()
|
||||
ax1 = fig1.add_subplot(111, aspect='equal')
|
||||
ax1.axes.xaxis.set_visible(False)
|
||||
ax1.axes.yaxis.set_visible(False)
|
||||
for node_name in self.node_pos:
|
||||
x, y, size_x, size_y = self.node_pos[node_name]
|
||||
ax1.add_patch(
|
||||
patches.Rectangle(
|
||||
(x/self.grid, y/self.grid), # (x,y)
|
||||
size_x/self.grid, # width
|
||||
size_y/self.grid, linewidth=1, edgecolor='k',
|
||||
)
|
||||
)
|
||||
fig1.savefig(file_path, dpi=90, bbox_inches='tight')
|
||||
plt.close()
|
||||
# WireMask
|
||||
def get_net_img(self, is_next_next = False):
|
||||
net_img = np.zeros((self.grid, self.grid))
|
||||
if not is_next_next:
|
||||
next_node_name = self.placedb.node_id_to_name[self.num_macro_placed]
|
||||
elif self.num_macro_placed + 1 < len(self.placedb.node_id_to_name):
|
||||
next_node_name = self.placedb.node_id_to_name[self.num_macro_placed + 1]
|
||||
else:
|
||||
return net_img
|
||||
|
||||
for net_name in self.placedb.node_to_net_dict[next_node_name]:
|
||||
if net_name in self.net_min_max_ord:
|
||||
delta_pin_x = round((self.placedb.node_info[next_node_name]['x']/2 + \
|
||||
self.placedb.net_info[net_name]["nodes"][next_node_name]["x_offset"])/self.ratio)
|
||||
delta_pin_y = round((self.placedb.node_info[next_node_name]['y']/2 + \
|
||||
self.placedb.net_info[net_name]["nodes"][next_node_name]["y_offset"])/self.ratio)
|
||||
start_x = self.net_min_max_ord[net_name]['min_x'] - delta_pin_x
|
||||
end_x = self.net_min_max_ord[net_name]['max_x'] - delta_pin_x
|
||||
start_y = self.net_min_max_ord[net_name]['min_y'] - delta_pin_y
|
||||
end_y = self.net_min_max_ord[net_name]['max_y'] - delta_pin_y
|
||||
start_x = min(start_x, self.grid)
|
||||
start_y = min(start_y, self.grid)
|
||||
if not 'weight' in self.placedb.net_info[net_name]:
|
||||
weight = 1.0
|
||||
else:
|
||||
weight = self.placedb.net_info[net_name]['weight']
|
||||
for i in range(0, start_x):
|
||||
net_img[i, :] += (start_x - i) * weight
|
||||
for i in range(end_x+1, self.grid):
|
||||
net_img[i, :] += (i- end_x) * weight
|
||||
for j in range(0, start_y):
|
||||
net_img[:, j] += (start_y - j) * weight
|
||||
for j in range(end_y+1, self.grid):
|
||||
net_img[:, j] += (j - start_y) * weight
|
||||
return net_img
|
||||
|
||||
def step(self, action):
|
||||
err_msg = f"{action!r} ({type(action)}) invalid"
|
||||
assert self.action_space.contains(action), err_msg
|
||||
|
||||
canvas = self.state[1: 1+self.grid*self.grid].reshape(self.grid, self.grid)
|
||||
mask = self.state[1+self.grid*self.grid*2: 1+self.grid*self.grid*3].reshape(self.grid, self.grid)
|
||||
reward = 0
|
||||
x = round(action // self.grid)
|
||||
y = round(action % self.grid)
|
||||
|
||||
if mask[x][y] == 1:
|
||||
reward += -200000
|
||||
|
||||
node_name = self.placedb.node_id_to_name[self.num_macro_placed]
|
||||
size_x = math.ceil(max(1, self.placedb.node_info[node_name]['x']/self.ratio))
|
||||
size_y = math.ceil(max(1, self.placedb.node_info[node_name]['y']/self.ratio))
|
||||
|
||||
assert abs(size_x - self.state[-2]*self.grid) < 1e-5
|
||||
assert abs(size_y - self.state[-1]*self.grid) < 1e-5
|
||||
|
||||
canvas[x : x+size_x, y : y+size_y] = 1.0
|
||||
canvas[x : x + size_x, y] = 0.5
|
||||
if y + size_y -1 < self.grid:
|
||||
canvas[x : x + size_x, max(0, y + size_y -1)] = 0.5
|
||||
canvas[x, y: y + size_y] = 0.5
|
||||
if x + size_x - 1 < self.grid:
|
||||
canvas[max(0, x+size_x-1), y: y + size_y] = 0.5
|
||||
self.node_pos[self.node_name_list[self.num_macro_placed]] = (x, y, size_x, size_y)
|
||||
|
||||
for net_name in self.placedb.node_to_net_dict[node_name]:
|
||||
self.net_placed_set[net_name].add(node_name)
|
||||
pin_x = round((x * self.ratio + self.placedb.node_info[node_name]['x']/2 + \
|
||||
self.placedb.net_info[net_name]["nodes"][node_name]["x_offset"])/self.ratio)
|
||||
pin_y = round((y * self.ratio + self.placedb.node_info[node_name]['y']/2 + \
|
||||
self.placedb.net_info[net_name]["nodes"][node_name]["y_offset"])/self.ratio)
|
||||
if net_name in self.net_min_max_ord:
|
||||
start_x = self.net_min_max_ord[net_name]['min_x']
|
||||
end_x = self.net_min_max_ord[net_name]['max_x']
|
||||
start_y = self.net_min_max_ord[net_name]['min_y']
|
||||
end_y = self.net_min_max_ord[net_name]['max_y']
|
||||
delta_x = end_x - start_x
|
||||
delta_y = end_y - start_y
|
||||
if delta_x > 0 or delta_y > 0:
|
||||
self.rudy[start_x : end_x +1, start_y: end_y +1] -= 1/(delta_x+1) + 1/(delta_y+1)
|
||||
weight = 1.0
|
||||
if 'weight' in self.placedb.net_info[net_name]:
|
||||
weight = self.placedb.net_info[net_name]['weight']
|
||||
|
||||
if pin_x > self.net_min_max_ord[net_name]['max_x']:
|
||||
reward += weight * (self.net_min_max_ord[net_name]['max_x'] - pin_x)
|
||||
self.net_min_max_ord[net_name]['max_x'] = pin_x
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][1] = pin_x / self.grid
|
||||
elif pin_x < self.net_min_max_ord[net_name]['min_x']:
|
||||
reward += weight * (pin_x - self.net_min_max_ord[net_name]['min_x'])
|
||||
self.net_min_max_ord[net_name]['min_x'] = pin_x
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][0] = pin_x / self.grid
|
||||
if pin_y > self.net_min_max_ord[net_name]['max_y']:
|
||||
reward += weight * (self.net_min_max_ord[net_name]['max_y'] - pin_y)
|
||||
self.net_min_max_ord[net_name]['max_y'] = pin_y
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][3] = pin_y / self.grid
|
||||
elif pin_y < self.net_min_max_ord[net_name]['min_y']:
|
||||
reward += weight * (pin_y - self.net_min_max_ord[net_name]['min_y'])
|
||||
self.net_min_max_ord[net_name]['min_y'] = pin_y
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][2] = pin_y / self.grid
|
||||
start_x = self.net_min_max_ord[net_name]['min_x']
|
||||
end_x = self.net_min_max_ord[net_name]['max_x']
|
||||
start_y = self.net_min_max_ord[net_name]['min_y']
|
||||
end_y = self.net_min_max_ord[net_name]['max_y']
|
||||
delta_x = end_x - start_x
|
||||
delta_y = end_y - start_y
|
||||
self.rudy[start_x : end_x +1, start_y: end_y +1] += 1/(delta_x+1) + 1/(delta_y+1)
|
||||
else:
|
||||
self.net_min_max_ord[net_name] = {}
|
||||
self.net_min_max_ord[net_name]['max_x'] = pin_x
|
||||
self.net_min_max_ord[net_name]['min_x'] = pin_x
|
||||
self.net_min_max_ord[net_name]['max_y'] = pin_y
|
||||
self.net_min_max_ord[net_name]['min_y'] = pin_y
|
||||
start_x = self.net_min_max_ord[net_name]['min_x']
|
||||
end_x = self.net_min_max_ord[net_name]['max_x']
|
||||
start_y = self.net_min_max_ord[net_name]['min_y']
|
||||
end_y = self.net_min_max_ord[net_name]['max_y']
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][1] = pin_x / self.grid
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][0] = pin_x / self.grid
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][3] = pin_y / self.grid
|
||||
self.net_fea[self.placedb.net_info[net_name]['id']][2] = pin_y / self.grid
|
||||
reward += 0
|
||||
|
||||
self.num_macro_placed += 1
|
||||
net_img = np.zeros((self.grid, self.grid))
|
||||
net_img_2 = np.zeros((self.grid, self.grid))
|
||||
|
||||
if self.num_macro_placed < self.placed_num_macro:
|
||||
net_img = self.get_net_img()
|
||||
net_img_2 = self.get_net_img(is_next_next= True)
|
||||
if net_img.max() >0 or net_img_2.max()>0:
|
||||
net_img /= max(net_img.max(), net_img_2.max())
|
||||
net_img_2 /= max(net_img.max(), net_img_2.max())
|
||||
|
||||
if self.node_x_max < x:
|
||||
self.node_x_max = x
|
||||
if self.node_x_min > x:
|
||||
self.node_x_min = x
|
||||
if self.node_y_max < y:
|
||||
self.node_y_max = y
|
||||
if self.node_y_min > y:
|
||||
self.node_y_min = y
|
||||
|
||||
if self.num_macro_placed == self.num_macro or \
|
||||
(self.placed_num_macro is not None and self.num_macro_placed == self.placed_num_macro):
|
||||
done = True
|
||||
else:
|
||||
done = False
|
||||
mask = np.ones((self.grid, self.grid))
|
||||
mask_2 = np.ones((self.grid, self.grid))
|
||||
if not done: # get next macro size and pre-mask the solution
|
||||
next_x = math.ceil(max(1, self.placedb.node_info[self.placedb.node_id_to_name[self.num_macro_placed]]['x']/self.ratio))
|
||||
next_y = math.ceil(max(1, self.placedb.node_info[self.placedb.node_id_to_name[self.num_macro_placed]]['y']/self.ratio))
|
||||
mask = self.get_mask(canvas, next_x, next_y)
|
||||
if self.num_macro_placed + 1 < self.placed_num_macro:
|
||||
next_x_2 = math.ceil(max(1, self.placedb.node_info[self.placedb.node_id_to_name[self.num_macro_placed+1]]['x']/self.ratio))
|
||||
next_y_2 = math.ceil(max(1, self.placedb.node_info[self.placedb.node_id_to_name[self.num_macro_placed+1]]['y']/self.ratio))
|
||||
mask_2 = self.get_mask(canvas, next_x_2, next_y_2)
|
||||
else:
|
||||
next_x = 0
|
||||
next_y = 0
|
||||
self.state = np.concatenate((np.array([self.num_macro_placed]), canvas.flatten(),
|
||||
net_img.flatten(), mask.flatten(), net_img_2.flatten(), mask_2.flatten(),
|
||||
np.array([next_x/self.grid, next_y/self.grid])), axis = 0)
|
||||
return self.state, reward, done, {"raw_reward": reward, "net_img": net_img, "mask": mask}
|
||||
|
||||
# PositionMask
|
||||
def get_mask(self, canvas, next_x, next_y):
|
||||
mask = np.zeros((self.grid, self.grid))
|
||||
for node_name in self.node_pos:
|
||||
startx = max(0, self.node_pos[node_name][0] - next_x + 1)
|
||||
starty = max(0, self.node_pos[node_name][1] - next_y + 1)
|
||||
endx = min(self.node_pos[node_name][0] + self.node_pos[node_name][2] - 1, self.grid - 1)
|
||||
endy = min(self.node_pos[node_name][1] + self.node_pos[node_name][3] - 1, self.grid - 1)
|
||||
mask[startx: endx + 1, starty : endy + 1] = 1
|
||||
mask[self.grid - next_x + 1:,:] = 1
|
||||
mask[:, self.grid - next_y + 1:] = 1
|
||||
return mask
|
||||
|
||||
|
||||
47
maskplace/prim.py
Normal file
47
maskplace/prim.py
Normal file
@ -0,0 +1,47 @@
|
||||
from itertools import combinations
|
||||
from heapq import *
|
||||
|
||||
def prim_real(vertexs_tmp, node_pos, net_info, ratio, node_info, port_info):# vertexs, edges,start='D'):
|
||||
vertexs = list(vertexs_tmp)
|
||||
if len(vertexs)<=1:
|
||||
return 0
|
||||
adjacent_dict = {}
|
||||
for node in vertexs:
|
||||
adjacent_dict[node] = []
|
||||
for node1, node2 in list(combinations(vertexs, 2)):
|
||||
if node1 in node_pos:
|
||||
pin_x_1 = node_pos[node1][0] * ratio + node_info[node1]["x"] / 2 + net_info[node1]["x_offset"] # )//ratio
|
||||
pin_y_1 = node_pos[node1][1] * ratio + node_info[node1]["y"] / 2 + net_info[node1]["y_offset"] # )//ratio
|
||||
else:
|
||||
pin_x_1 = port_info[node1]['x']
|
||||
pin_y_1 = port_info[node1]['y']
|
||||
if node2 in node_pos:
|
||||
pin_x_2 = node_pos[node2][0] * ratio + node_info[node2]["x"] / 2 + net_info[node2]["x_offset"] # )//ratio
|
||||
pin_y_2 = node_pos[node2][1] * ratio + node_info[node2]["y"] / 2 + net_info[node2]["y_offset"] # )//ratio
|
||||
else:
|
||||
pin_x_2 = port_info[node2]['x']
|
||||
pin_y_2 = port_info[node2]['y']
|
||||
weight = abs(pin_x_1-pin_x_2) + \
|
||||
abs(pin_y_1-pin_y_2)
|
||||
adjacent_dict[node1].append((weight, node1, node2))
|
||||
adjacent_dict[node2].append((weight, node2, node1))
|
||||
|
||||
start = vertexs[0]
|
||||
minu_tree = []
|
||||
visited = set()
|
||||
visited.add(start)
|
||||
adjacent_vertexs_edges = adjacent_dict[start]
|
||||
heapify(adjacent_vertexs_edges)
|
||||
cost = 0
|
||||
cnt = 0
|
||||
while cnt < len(vertexs)-1:
|
||||
weight, v1, v2 = heappop(adjacent_vertexs_edges)
|
||||
if v2 not in visited:
|
||||
visited.add(v2)
|
||||
minu_tree.append((weight, v1, v2))
|
||||
cost += weight
|
||||
cnt += 1
|
||||
for next_edge in adjacent_dict[v2]:
|
||||
if next_edge[2] not in visited:
|
||||
heappush(adjacent_vertexs_edges, next_edge)
|
||||
return cost
|
||||
Loading…
Reference in New Issue
Block a user