Skip to content

Commit 1c9fcb8

Browse files
author
史昊林
committed
feat: Qwen2模型实现
1 parent 958c445 commit 1c9fcb8

13 files changed

Lines changed: 984 additions & 153 deletions

File tree

include/llaisys/models/qwen2.h

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -29,14 +29,16 @@ __C {
2929
llaisysTensor_t *mlp_down_w;
3030
};
3131

32-
struct LlaisysQwen2Model;
33-
32+
struct LlaisysQwen2Model {
33+
struct LlaisysQwen2Meta *meta;
34+
llaisysDeviceType_t device;
35+
int ndevice;
36+
int *device_ids;
37+
struct LlaisysQwen2Weights *weights;
38+
};
3439
__export struct LlaisysQwen2Model *llaisysQwen2ModelCreate(const LlaisysQwen2Meta *meta, llaisysDeviceType_t device, int *device_ids, int ndevice);
35-
3640
__export void llaisysQwen2ModelDestroy(struct LlaisysQwen2Model * model);
37-
3841
__export struct LlaisysQwen2Weights *llaisysQwen2ModelWeights(struct LlaisysQwen2Model * model);
39-
40-
__export int64_t llaisysQwen2ModelInfer(struct LlaisysQwen2Model * model, int64_t * token_ids, size_t ntoken);
42+
__export int64_t llaisysQwen2ModelInfer(struct LlaisysQwen2Model * model, int64_t * token_ids, size_t ntoken, llaisysTensor_t *kcache, llaisysTensor_t *vcache, size_t past_len);
4143
}
4244
#endif // LLAISYS_MODELS_QWEN2_H

include/llaisys/tensor.h

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -54,6 +54,11 @@ __C {
5454
size_t * shape,
5555
size_t ndim);
5656

57+
__export llaisysTensor_t tensorReshape(
58+
llaisysTensor_t tensor,
59+
size_t * shape,
60+
size_t ndim);
61+
5762
__export llaisysTensor_t tensorPermute(
5863
llaisysTensor_t tensor,
5964
size_t * order);
Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
from .qwen2 import load_qwen2, LlaisysQwen2Meta, LlaisysQwen2Weights
Lines changed: 65 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,65 @@
1+
from ctypes import c_int, c_int64, c_size_t
2+
from ..llaisys_types import llaisysDeviceType_t, llaisysDataType_t
3+
from ..tensor import llaisysTensor_t
4+
import ctypes
5+
6+
class LlaisysQwen2Meta(ctypes.Structure):
7+
_fields_ = [
8+
("dtype", llaisysDataType_t),
9+
("nlayer", ctypes.c_size_t),
10+
("hs", ctypes.c_size_t),
11+
("nh", ctypes.c_size_t),
12+
("nkvh", ctypes.c_size_t),
13+
("dh", ctypes.c_size_t),
14+
("di", ctypes.c_size_t),
15+
("maxseq", ctypes.c_size_t),
16+
("voc", ctypes.c_size_t),
17+
("epsilon", ctypes.c_float),
18+
("theta", ctypes.c_float),
19+
("end_token", ctypes.c_int64)
20+
]
21+
22+
class LlaisysQwen2Weights(ctypes.Structure):
23+
_fields_ = [
24+
("in_embed", llaisysTensor_t),
25+
("out_embed", llaisysTensor_t),
26+
("out_norm_w", llaisysTensor_t),
27+
28+
("attn_norm_w", ctypes.POINTER(llaisysTensor_t)),
29+
("attn_q_w", ctypes.POINTER(llaisysTensor_t)),
30+
("attn_q_b", ctypes.POINTER(llaisysTensor_t)),
31+
("attn_k_w", ctypes.POINTER(llaisysTensor_t)),
32+
("attn_k_b", ctypes.POINTER(llaisysTensor_t)),
33+
("attn_v_w", ctypes.POINTER(llaisysTensor_t)),
34+
("attn_v_b", ctypes.POINTER(llaisysTensor_t)),
35+
("attn_o_w", ctypes.POINTER(llaisysTensor_t)),
36+
37+
("mlp_norm_w", ctypes.POINTER(llaisysTensor_t)),
38+
("mlp_gate_w", ctypes.POINTER(llaisysTensor_t)),
39+
("mlp_up_w", ctypes.POINTER(llaisysTensor_t)),
40+
("mlp_down_w", ctypes.POINTER(llaisysTensor_t)),
41+
]
42+
43+
class LlaisysQwen2Model(ctypes.Structure):
44+
_fields_ = [
45+
("meta", ctypes.POINTER(LlaisysQwen2Meta)),
46+
("device", ctypes.c_int), # llaisysDeviceType_t
47+
("ndevice", ctypes.c_int),
48+
("device_ids", ctypes.POINTER(ctypes.c_int)),
49+
("weights", ctypes.POINTER(LlaisysQwen2Weights)),
50+
]
51+
52+
# Load shared library
53+
def load_qwen2(lib):
54+
# Declare API function prototypes
55+
lib.llaisysQwen2ModelCreate.argtypes = [ctypes.POINTER(LlaisysQwen2Meta), llaisysDeviceType_t, ctypes.POINTER(c_int), c_int]
56+
lib.llaisysQwen2ModelCreate.restype = ctypes.POINTER(LlaisysQwen2Model)
57+
58+
lib.llaisysQwen2ModelDestroy.argtypes = [ctypes.POINTER(LlaisysQwen2Model)]
59+
lib.llaisysQwen2ModelDestroy.restype = None
60+
61+
lib.llaisysQwen2ModelWeights.argtypes = [ctypes.POINTER(LlaisysQwen2Model)]
62+
lib.llaisysQwen2ModelWeights.restype = ctypes.POINTER(LlaisysQwen2Weights)
63+
64+
lib.llaisysQwen2ModelInfer.argtypes = [ctypes.POINTER(LlaisysQwen2Model), ctypes.POINTER(c_int64), c_size_t, ctypes.POINTER(llaisysTensor_t), ctypes.POINTER(llaisysTensor_t), c_size_t]
65+
lib.llaisysQwen2ModelInfer.restype = c_int64

python/llaisys/libllaisys/tensor.py

Lines changed: 14 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -6,73 +6,56 @@
66

77

88
def load_tensor(lib):
9-
lib.tensorCreate.argtypes = [
10-
POINTER(c_size_t), # shape
11-
c_size_t, # ndim
12-
llaisysDataType_t, # dtype
13-
llaisysDeviceType_t, # device_type
14-
c_int, # device_id
15-
]
9+
"""Configure tensor function signatures for the C library."""
10+
11+
# Core tensor lifecycle functions
12+
lib.tensorCreate.argtypes = [POINTER(c_size_t), c_size_t, llaisysDataType_t, llaisysDeviceType_t, c_int]
1613
lib.tensorCreate.restype = llaisysTensor_t
1714

18-
# Function: tensorDestroy
1915
lib.tensorDestroy.argtypes = [llaisysTensor_t]
2016
lib.tensorDestroy.restype = None
2117

22-
# Function: tensorGetData
18+
# Tensor property accessors
2319
lib.tensorGetData.argtypes = [llaisysTensor_t]
2420
lib.tensorGetData.restype = c_void_p
2521

26-
# Function: tensorGetNdim
2722
lib.tensorGetNdim.argtypes = [llaisysTensor_t]
2823
lib.tensorGetNdim.restype = c_size_t
2924

30-
# Function: tensorGetShape
3125
lib.tensorGetShape.argtypes = [llaisysTensor_t, POINTER(c_size_t)]
3226
lib.tensorGetShape.restype = None
3327

34-
# Function: tensorGetStrides
3528
lib.tensorGetStrides.argtypes = [llaisysTensor_t, POINTER(c_ssize_t)]
3629
lib.tensorGetStrides.restype = None
3730

38-
# Function: tensorGetDataType
3931
lib.tensorGetDataType.argtypes = [llaisysTensor_t]
4032
lib.tensorGetDataType.restype = llaisysDataType_t
4133

42-
# Function: tensorGetDeviceType
4334
lib.tensorGetDeviceType.argtypes = [llaisysTensor_t]
4435
lib.tensorGetDeviceType.restype = llaisysDeviceType_t
4536

46-
# Function: tensorGetDeviceId
4737
lib.tensorGetDeviceId.argtypes = [llaisysTensor_t]
4838
lib.tensorGetDeviceId.restype = c_int
4939

50-
# Function: tensorDebug
51-
lib.tensorDebug.argtypes = [llaisysTensor_t]
52-
lib.tensorDebug.restype = None
53-
54-
# Function: tensorIsContiguous
5540
lib.tensorIsContiguous.argtypes = [llaisysTensor_t]
5641
lib.tensorIsContiguous.restype = c_uint8
5742

58-
# Function: tensorLoad
43+
# Data manipulation functions
5944
lib.tensorLoad.argtypes = [llaisysTensor_t, c_void_p]
6045
lib.tensorLoad.restype = None
6146

62-
# Function: tensorView(llaisysTensor_t tensor, size_t *shape);
47+
lib.tensorDebug.argtypes = [llaisysTensor_t]
48+
lib.tensorDebug.restype = None
49+
50+
# Tensor transformation functions
6351
lib.tensorView.argtypes = [llaisysTensor_t, POINTER(c_size_t), c_size_t]
6452
lib.tensorView.restype = llaisysTensor_t
6553

66-
# Function: tensorPermute(llaisysTensor_t tensor, size_t *order);
54+
lib.tensorReshape.argtypes = [llaisysTensor_t, POINTER(c_size_t), c_size_t]
55+
lib.tensorReshape.restype = llaisysTensor_t
56+
6757
lib.tensorPermute.argtypes = [llaisysTensor_t, POINTER(c_size_t)]
6858
lib.tensorPermute.restype = llaisysTensor_t
6959

70-
# Function: tensorSlice(llaisysTensor_t tensor,
71-
# size_t dim, size_t start, size_t end);
72-
lib.tensorSlice.argtypes = [
73-
llaisysTensor_t, # tensor handle
74-
c_size_t, # dim : which axis to slice
75-
c_size_t, # start: inclusive
76-
c_size_t, # end : exclusive
77-
]
60+
lib.tensorSlice.argtypes = [llaisysTensor_t, c_size_t, c_size_t, c_size_t]
7861
lib.tensorSlice.restype = llaisysTensor_t

0 commit comments

Comments
 (0)