QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 3438|回复: 0
打印 上一主题 下一主题

[其他资源] K-近邻算法分类和回归

[复制链接]
字体大小: 正常 放大
杨利霞        

5273

主题

82

听众

17万

积分

  • TA的每日心情
    开心
    2021-8-11 17:59
  • 签到天数: 17 天

    [LV.4]偶尔看看III

    网络挑战赛参赛者

    网络挑战赛参赛者

    自我介绍
    本人女,毕业于内蒙古科技大学,担任文职专业,毕业专业英语。

    群组2018美赛大象算法课程

    群组2018美赛护航培训课程

    群组2019年 数学中国站长建

    群组2019年数据分析师课程

    群组2018年大象老师国赛优

    跳转到指定楼层
    1#
    发表于 2022-9-5 15:43 |只看该作者 |倒序浏览
    |招呼Ta 关注Ta
    . i  K' J) N. u! S4 r
    K-近邻算法分类和回归& E, U2 C4 T7 t5 n+ Q
    K近邻算法的主要思想是用离测试集数据点最近的训练集点(称为其邻居)的输出来估计测试集数据点的输出,参数K代表用多少个邻居来估计。超参数K通常设置为奇数来防止平局现象。0 c# P$ I" `+ {* D7 b
    4 c$ E% G; [+ t
    其中对邻居的判定:我们可以用欧几里得距离来衡量距离来确定其K个邻居。
    ! B  x$ G* ~+ z  V# e$ f
    : p# h0 z- H4 _8 _3 d/ E' v# dK近邻算法是一种惰性学习和非参数模型。当训练数据数量庞大,同时你对响应变量和解释变量之间的关系所知甚少时,非参数模型会非常有用。KNN 模型只基于一个假设:互相接近的实例拥有类似的响应变量值。非参数模型提供的灵活性并不总是可取的,当训练数据很缺乏或者你对响应变量和解释变量之间的关系有所了解时,对响应变量和解释变量之间关系做假设的模型就很有用。' L3 r+ c0 G  ]4 [5 h5 c8 N
    : H9 P) d0 e4 d
    KNN模型分类:
    % {. b% `1 p5 ~1 b4 _下面我们看一个分类的例子和代码实现来了解一下K近邻算法:. ]" }: C* N" K8 |. ?% b
    . m. g+ T9 b. ~( x$ n% Y7 ?; x
    # l8 ?: a! T* _4 u; x4 q2 Z
    0 F' b3 z. m$ x( X
    上表是我们的训练集,下面先对数据进行可视化
    / H4 K, P" B, Z) B$ Z  a, I
    & I! l# r4 m& t9 x5 Q" ximport numpy as np
    - B% W8 I; {( M+ l  b) |" @/ E9 U' U& afrom matplotlib import pyplot as plt. ^* u9 K9 x) _; Q3 h, B5 x+ v
    import sklearn5 p! V- d0 J' `! i+ |/ i) C
    7 f: k+ N/ m; E  f0 [; O8 S
    X_train = np.array([ # 身高体重4 A/ t7 A' |/ c5 P) C" h* \
        [158, 64],
    9 X5 s0 q3 m* j5 K/ W    [170, 86],
    * g9 Y' o& d5 W- ]# j) k8 R( T    [183, 84],) z1 P2 D; T- X
        [191, 80],% d1 b) E$ q3 p* {% n
        [155, 49],
    , J9 z( Z4 x" A2 j0 J5 `    [163, 59],1 U) X- n1 P  u* e3 [: J$ P
        [180, 67],2 _: |8 o7 n% e' T# g: F9 `- v. S
        [158, 54],# d! o% w' Z5 f5 _% |
        [170, 67]])0 c0 F( u% R7 \( S  z1 A
    y_train = ['male']*4 + ['female']*5 # 性别; ^( J  ?  u8 L) a. s5 H8 u$ h

    8 E- }) x) n0 ]+ H0 ~+ O0 `#绘制图像  ]/ y% b0 A' e4 a3 d6 l
    plt.figure()
    - s, Q$ r: y1 W) Q' ?plt.title('Human Height and Weights by Sex')
    ' X  C7 y* y3 ^% I, z5 Gplt.xlabel('Height in cm')  m  g7 D( ]2 m. ?
    plt.ylabel('Weight in kg')9 Z3 K5 v# I2 I4 j" h% h
    for i, x in enumerate(X_train):$ s- s, L$ x% G! o! M
        plt.scatter(x[0],x[1],c='k',marker='x' if y_train == 'male' else 'D')
    6 v& n; W* e8 Aplt.grid()
    7 K* U  r- N7 F  _* Lplt.show()- I# Z7 n' H3 g( i
    2 m1 `, |# W7 [
    结果:
    : k2 e' Z" ]8 q4 ~7 c
    " q6 v- z& P% Q% g: ~$ N  [, @% T  x, Y/ H& U0 f2 ?

    5 E- b2 O% ~7 E; w; R% n) m& c 我们使用欧几里得距离公式来衡量距离:
    5 o( G2 ]8 k  d+ l! R" b8 p) f; v- R4 M. t/ Q) R" I( v; s% r, Y
    9 \& ~. b: _- W/ Z. w
    # N/ v5 _0 I8 U: F4 a$ Z/ m
    4 t9 m7 O8 _  D' k% E' G8 N" e
    4 M  S; `8 _$ c3 T/ R# C
    我们设置参数K=3,来寻找3个距离最近的训练实例
    2 j2 }. T$ j5 p9 ~+ V1 N
    $ `* a7 H: _8 {8 `+ K! e+ w下面代码实现K近邻算法进行分类:
    0 l' v' ~: j. O7 C0 G' e% |
    ; J( T* q1 E  e! R8 u: Mx = np.array([[155, 70]])
    3 Z: Q1 z9 ^, u# Edistances = np.sqrt(np.sum((X_train - x)**2, axis=1)) # 计算距离3 u' i+ O4 s( w7 |
    ( u$ E' P2 U" n. F6 W, ?% q" K
    nearest_neighbor_indices = distances.argsort()[:3] # 找出前三个距离最小的下标
    6 E! z. v7 v( @% n  ~$ j) _nearest_neighbor_genders = np.take(y_train, nearest_neighbor_indices) # 得到下标对应的标签$ F! a) ]. W: _6 T" D# f* U
    0 w3 l$ o  J/ g
    from collections import Counter) {& y- l$ y4 A2 r! p
    b = Counter(np.take(y_train, distances.argsort()[:3])) #得到三个结果标签中最频繁的标签得到结果female
    4 A0 q- |0 B. Q9 z9 a" S$ c& P( r
    & q; @5 u1 }! @& Y# m2 tprint(b.most_common(1)[0][0]) # female# b8 z8 n! ]# d
    因此,从上述代码可以得到K近邻算法进行分类就是找到离样本点最近的K个实例,再取K个实例的标签中出现次数最多的那个作为我们的结果。
    2 @1 s% t4 ^1 z/ O. e  J
    ( g5 g% x) V* N' g0 X/ \上述K近邻算法在scikit-learn中也有对应的函数:
    ( d: ^, s9 p* w" W
    $ X) V1 n) Y; V! m; Efrom sklearn.preprocessing import LabelBinarizer( \5 x  }$ ~$ w: ]: V3 ?
    from sklearn.neighbors import KNeighborsClassifier1 A0 V. s" b. D3 O+ J

    8 Z2 B! {8 D5 v& {& Vlb = LabelBinarizer() # 创建将标签二值数值化的类实例# ]! _" S  h/ M. v8 g* T2 Y. Q
    y_train_binarized = lb.fit_transform(y_train) # 将标签二值数值化
    # ^0 h* A5 z8 r; f4 F5 rprint(y_train_binarized) ! O( I( A- L1 I  O- W% o
    - B) ], G0 M9 h! Q; ~8 f
    K = 3; J/ X' W/ r$ ?8 E3 M/ A( [; Y2 c
    clf = KNeighborsClassifier(n_neighbors=K) # 创建K近邻分类器实例
    / l. Z) i2 s  i$ [( W* Z! o/ Mclf.fit(X_train, y_train_binarized.reshape(-1, 1)) # 对训练集进行训练5 ]! n& m7 {; G( k* b" D4 P6 J
    prediction_binarized = clf.predict(np.array([155, 70]).reshape(1,-1))[0] # 对测试样本点进行预测
    - Q& d/ {* F" `prediction_label = lb.inverse_transform(prediction_binarized) # 将预测结果从数字转换为标签
    1 E; b, S$ y5 `8 g- B# {print(prediction_label) # array(['female'], dtype='<U6')
    ( h4 k, K0 ~5 |% i) ^1 R, j  J# v1 rKNN回归:/ p3 ^# w( |9 Q6 g# ]' L( E. o9 {3 P
    K近邻算法进行回归和K近邻思想一致,只不过在得到了K个邻居后,分类是取邻居中出现次数最多的那个,而回归是取其他的操作来预测输出值(比如取平均)
    1 ~5 @- g1 F4 N" `* v" u
    ' G& [- W. A$ y" n对应的代码在scikit-learn中其实也很简单
    $ }2 A! u3 `, n; g- d0 k0 G" Z4 b8 F2 ]: A! I
    from sklearn.neighbors import KNeighborsRegressor7 n! z: G; u) g* C) Y/ {; l
    K = 3
    - e6 r6 }9 z+ q4 h1 Uclf = KNeighborsRegressor(n_neighbors=K)
    8 G0 x( [) H+ h! pclf.fit(X_train, y_train), l; h/ Z/ k8 Q$ F! @
    predictions = clf.predict(X_test): Q  ^$ l  ~) F& h0 w: f2 E2 m
    特征缩放! B. S& Q3 r  b2 L6 R
    下面我们谈谈一个提升算法精确度的小细节。假设还是上面的数据,我们现在要做回归,给定身高和性别标签来预测体重。如果我们的训练数据集包含一个身高170cm的男性和身高160cm的女性。如果我们的测试集数据为身高为164cm的男性,你觉得其预测结果会接近170cm的男性还是身高160cm的女性呢?我们可能相信测试实例更接近男性实例,因为对预测体重来说,性别差异可能会比 6cm 的身高差距更重要。但是如果我们以毫米为单位表示身高,测试实例更接近于身高1600mm 的女性。如果我们以米为单位表示身高,测试实例更接近于身高 1.7m 的男性。(记住我们以欧几里得距离来衡量)3 n- v2 e( {3 v, W/ c1 H( q& U
    # c- W( q6 ?2 Y2 J" u! l4 E  i( e9 f: C' s. R
    因此,我们的特征缩放的作用就出来了(其实就相当于深度学习对数据集预处理中的Normalize)
    2 r" X" d) p3 k# ^; Y8 z& L, F( A" c. O- X! D
    将所有实例特征值减去均值来将其居中。其次将每个实例特征值除以特征的标准差对其进行缩放。均值为 0,方差为 1 的数据称为标准化数据。4 a7 a3 i# ~2 ^" P' r1 ]# W
    ————————————————
    & P5 C/ x0 h' l8 a' R7 }/ s版权声明:本文为CSDN博主「王大队长」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。$ g! |. o, j4 o$ D+ J& a: b
    原文链接:https://blog.csdn.net/qq_55621259/article/details/1266955496 |9 z! a5 Q; ]
    " i. ]; h/ [# J6 F' h( L

    . |! E3 |  c7 |5 P: d) g5 \
    zan
    转播转播0 分享淘帖0 分享分享0 收藏收藏0 支持支持0 反对反对0 微信微信
    您需要登录后才可以回帖 登录 | 注册地址

    qq
    收缩
    • 电话咨询

    • 04714969085
    fastpost

    关于我们| 联系我们| 诚征英才| 对外合作| 产品服务| QQ

    手机版|Archiver| |繁體中文 手机客户端  

    蒙公网安备 15010502000194号

    Powered by Discuz! X2.5   © 2001-2013 数学建模网-数学中国 ( 蒙ICP备14002410号-3 蒙BBS备-0002号 )     论坛法律顾问:王兆丰

    GMT+8, 2026-7-29 04:18 , Processed in 0.434750 second(s), 51 queries .

    回顶部